chore: import upstream snapshot with attribution
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
This commit is contained in:
@@ -0,0 +1,109 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.lora.utils import LoRABatchInfo
|
||||
|
||||
from .graph_lora_ops import (
|
||||
sgemm_lora_a_embedding_graph_fwd,
|
||||
sgemm_lora_a_graph_fwd,
|
||||
sgemm_lora_b_graph_fwd,
|
||||
)
|
||||
from .lora_ops import sgemm_lora_a_embedding_fwd as sgemm_lora_a_embedding_control_fwd
|
||||
from .lora_ops import sgemm_lora_a_fwd as sgemm_lora_a_control_fwd
|
||||
from .lora_ops import sgemm_lora_b_fwd as sgemm_lora_b_control_fwd
|
||||
|
||||
|
||||
def sgemm_lora_a_embedding_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
batch_info: LoRABatchInfo,
|
||||
vocab_size: int,
|
||||
) -> torch.Tensor:
|
||||
output: torch.Tensor
|
||||
if batch_info.use_cuda_graph:
|
||||
output = sgemm_lora_a_embedding_graph_fwd(
|
||||
inputs,
|
||||
weights,
|
||||
batch_info.weight_indices,
|
||||
batch_info.seg_lens,
|
||||
batch_info.scalings,
|
||||
vocab_size,
|
||||
)
|
||||
else:
|
||||
output = sgemm_lora_a_embedding_control_fwd(
|
||||
inputs,
|
||||
weights,
|
||||
batch_info.weight_indices_cpu,
|
||||
batch_info.seg_lens_cpu,
|
||||
batch_info.lora_ranks_cpu,
|
||||
batch_info.scalings_cpu,
|
||||
vocab_size,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def sgemm_lora_a_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
batch_info: LoRABatchInfo,
|
||||
num_slices: int = 1,
|
||||
) -> torch.Tensor:
|
||||
output: torch.Tensor
|
||||
if batch_info.use_cuda_graph:
|
||||
output = sgemm_lora_a_graph_fwd(
|
||||
inputs,
|
||||
weights,
|
||||
batch_info.weight_indices,
|
||||
batch_info.seg_lens,
|
||||
batch_info.scalings,
|
||||
num_slices,
|
||||
)
|
||||
else:
|
||||
output = sgemm_lora_a_control_fwd(
|
||||
inputs,
|
||||
weights,
|
||||
batch_info.weight_indices_cpu,
|
||||
batch_info.seg_lens_cpu,
|
||||
batch_info.lora_ranks_cpu,
|
||||
batch_info.scalings_cpu,
|
||||
num_slices,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def sgemm_lora_b_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
batch_info: LoRABatchInfo,
|
||||
slice_offsets: torch.Tensor,
|
||||
base_output: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
output: torch.Tensor
|
||||
if batch_info.use_cuda_graph:
|
||||
output = sgemm_lora_b_graph_fwd(
|
||||
inputs,
|
||||
weights,
|
||||
batch_info.weight_indices,
|
||||
batch_info.seg_lens,
|
||||
slice_offsets,
|
||||
base_output,
|
||||
)
|
||||
else:
|
||||
output = sgemm_lora_b_control_fwd(
|
||||
inputs,
|
||||
weights,
|
||||
batch_info.weight_indices_cpu,
|
||||
batch_info.seg_lens_cpu,
|
||||
batch_info.lora_ranks_cpu,
|
||||
slice_offsets,
|
||||
base_output,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
__all__ = [
|
||||
"sgemm_lora_a_embedding_fwd",
|
||||
"sgemm_lora_a_fwd",
|
||||
"sgemm_lora_b_fwd",
|
||||
]
|
||||
@@ -0,0 +1,120 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def sgemm_lora_a_embedding_graph_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
weight_indices: torch.Tensor,
|
||||
seg_len_tensor: torch.Tensor,
|
||||
scaling_tensor: torch.Tensor,
|
||||
vocab_size: int,
|
||||
) -> torch.Tensor:
|
||||
total_seq_len = inputs.shape[0]
|
||||
if weights.numel() == 0:
|
||||
return torch.zeros(total_seq_len, 0, dtype=weights.dtype, device=weights.device)
|
||||
|
||||
num_loras, max_rank, _ = weights.shape
|
||||
|
||||
output = torch.zeros(
|
||||
total_seq_len, max_rank, dtype=weights.dtype, device=weights.device
|
||||
)
|
||||
|
||||
for lora_idx in range(num_loras):
|
||||
|
||||
batch_token_mask = weight_indices[:total_seq_len] == lora_idx
|
||||
|
||||
x_seq = torch.where(batch_token_mask, inputs, 0)
|
||||
w_seq = weights[lora_idx]
|
||||
|
||||
output.add_(
|
||||
scaling_tensor[lora_idx]
|
||||
* torch.where(
|
||||
batch_token_mask.unsqueeze(1), F.embedding(x_seq, w_seq.t()), 0
|
||||
)
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def sgemm_lora_a_graph_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
weight_indices: torch.Tensor,
|
||||
seg_len_tensor: torch.Tensor,
|
||||
scaling_tensor: torch.Tensor,
|
||||
num_slices: int = 1,
|
||||
) -> torch.Tensor:
|
||||
total_seq_len, input_dim = inputs.shape
|
||||
if weights.numel() == 0:
|
||||
return torch.zeros(total_seq_len, 0, dtype=inputs.dtype, device=inputs.device)
|
||||
|
||||
num_loras, weight_out_dim, _ = weights.shape
|
||||
max_rank = weight_out_dim // num_slices
|
||||
|
||||
output = torch.zeros(
|
||||
total_seq_len, num_slices * max_rank, dtype=inputs.dtype, device=inputs.device
|
||||
)
|
||||
|
||||
for lora_idx in range(num_loras):
|
||||
|
||||
batch_token_mask = (weight_indices[:total_seq_len] == lora_idx).unsqueeze(1)
|
||||
|
||||
x_seq = torch.where(batch_token_mask, inputs, 0)
|
||||
w_seq = weights[lora_idx]
|
||||
|
||||
output.add_(scaling_tensor[lora_idx] * torch.mm(x_seq, w_seq.t()))
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def sgemm_lora_b_graph_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
weight_indices: torch.Tensor,
|
||||
seg_len_tensor: torch.Tensor,
|
||||
slice_offsets: torch.Tensor,
|
||||
base_output: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
total_seq_len, input_dim = inputs.shape
|
||||
num_loras, weight_out_dim, _ = weights.shape
|
||||
total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0
|
||||
|
||||
if weights.numel() == 0:
|
||||
return torch.zeros(
|
||||
total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device
|
||||
)
|
||||
|
||||
num_slices = len(slice_offsets) - 1
|
||||
max_rank = input_dim // num_slices
|
||||
|
||||
if base_output is not None:
|
||||
output = base_output
|
||||
else:
|
||||
output = torch.zeros(
|
||||
total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device
|
||||
)
|
||||
|
||||
for lora_idx in range(num_loras):
|
||||
|
||||
batch_token_mask = (weight_indices[:total_seq_len] == lora_idx).unsqueeze(1)
|
||||
inputs_masked = torch.where(batch_token_mask, inputs, 0)
|
||||
|
||||
for slice_idx in range(num_slices):
|
||||
slice_start_input = slice_idx * max_rank
|
||||
slice_end_input = (slice_idx + 1) * max_rank
|
||||
|
||||
slice_start_output = slice_offsets[slice_idx]
|
||||
slice_end_output = slice_offsets[slice_idx + 1]
|
||||
|
||||
x_slice = inputs_masked[..., slice_start_input:slice_end_input]
|
||||
w_slice = weights[
|
||||
lora_idx, slice_start_output:slice_end_output
|
||||
] # (slice_dim, max_rank)
|
||||
output[..., slice_start_output:slice_end_output].add_(
|
||||
torch.mm(x_slice, w_slice.t())
|
||||
)
|
||||
|
||||
return output
|
||||
@@ -0,0 +1,146 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def sgemm_lora_a_embedding_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
weight_indices: torch.Tensor,
|
||||
seg_len_tensor: torch.Tensor,
|
||||
lora_ranks: torch.Tensor,
|
||||
scaling_tensor: torch.Tensor,
|
||||
vocab_size: int,
|
||||
) -> torch.Tensor:
|
||||
total_seq_len = inputs.shape[0]
|
||||
if weights.numel() == 0:
|
||||
return torch.zeros(total_seq_len, 0, dtype=weights.dtype, device=weights.device)
|
||||
|
||||
num_loras, max_rank, _ = weights.shape
|
||||
|
||||
output = torch.zeros(
|
||||
total_seq_len, max_rank, dtype=weights.dtype, device=weights.device
|
||||
)
|
||||
|
||||
token_offset = 0
|
||||
for lora_idx, seq_len in zip(weight_indices, seg_len_tensor):
|
||||
if seq_len == 0:
|
||||
continue
|
||||
|
||||
rank = lora_ranks[lora_idx]
|
||||
if rank > 0:
|
||||
|
||||
x_seq = inputs[token_offset : token_offset + seq_len]
|
||||
w_seq = weights[lora_idx, :rank]
|
||||
|
||||
result = torch.nn.functional.embedding(x_seq, w_seq.T)
|
||||
output[token_offset : token_offset + seq_len, :rank] = (
|
||||
scaling_tensor[lora_idx].item() * result
|
||||
)
|
||||
|
||||
token_offset += seq_len
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def sgemm_lora_a_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
weight_indices: torch.Tensor,
|
||||
seg_len_tensor: torch.Tensor,
|
||||
lora_ranks: torch.Tensor,
|
||||
scaling_tensor: torch.Tensor,
|
||||
num_slices: int = 1,
|
||||
) -> torch.Tensor:
|
||||
total_seq_len, input_dim = inputs.shape
|
||||
if weights.numel() == 0:
|
||||
return torch.zeros(total_seq_len, 0, dtype=inputs.dtype, device=inputs.device)
|
||||
|
||||
num_loras, weight_out_dim, _ = weights.shape
|
||||
max_rank = weight_out_dim // num_slices
|
||||
|
||||
output = torch.zeros(
|
||||
total_seq_len, num_slices * max_rank, dtype=inputs.dtype, device=inputs.device
|
||||
)
|
||||
|
||||
token_offset = 0
|
||||
for lora_idx, seq_len in zip(weight_indices, seg_len_tensor):
|
||||
if seq_len == 0:
|
||||
continue
|
||||
|
||||
rank = lora_ranks[lora_idx]
|
||||
if rank > 0:
|
||||
|
||||
x_seq = inputs[token_offset : token_offset + seq_len]
|
||||
w_seq = weights[lora_idx, : num_slices * rank]
|
||||
|
||||
output[token_offset : token_offset + seq_len, : num_slices * rank].addmm_(
|
||||
x_seq,
|
||||
w_seq.T,
|
||||
beta=0,
|
||||
alpha=scaling_tensor[lora_idx].item(),
|
||||
)
|
||||
|
||||
token_offset += seq_len
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def sgemm_lora_b_fwd(
|
||||
inputs: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
weight_indices: torch.Tensor,
|
||||
seg_len_tensor: torch.Tensor,
|
||||
lora_ranks: torch.Tensor,
|
||||
slice_offsets: torch.Tensor,
|
||||
base_output: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
total_seq_len, _ = inputs.shape
|
||||
num_loras, weight_out_dim, _ = weights.shape
|
||||
total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0
|
||||
|
||||
if weights.numel() == 0:
|
||||
return torch.zeros(
|
||||
total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device
|
||||
)
|
||||
|
||||
num_slices = len(slice_offsets) - 1
|
||||
|
||||
if base_output is not None:
|
||||
output = base_output
|
||||
else:
|
||||
output = torch.zeros(
|
||||
total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device
|
||||
)
|
||||
|
||||
token_offset = 0
|
||||
for lora_idx, seq_len in zip(weight_indices, seg_len_tensor):
|
||||
if seq_len == 0:
|
||||
continue
|
||||
|
||||
rank = lora_ranks[lora_idx]
|
||||
if rank > 0:
|
||||
|
||||
for slice_idx in range(num_slices):
|
||||
slice_start_input = slice_idx * rank
|
||||
slice_end_input = (slice_idx + 1) * rank
|
||||
|
||||
slice_start_output = slice_offsets[slice_idx]
|
||||
slice_end_output = slice_offsets[slice_idx + 1]
|
||||
|
||||
x_slice = inputs[
|
||||
token_offset : token_offset + seq_len,
|
||||
slice_start_input:slice_end_input,
|
||||
] # (seq_len, rank)
|
||||
w_slice = weights[
|
||||
lora_idx, slice_start_output:slice_end_output, :rank
|
||||
] # (slice_dim, rank)
|
||||
|
||||
output[
|
||||
token_offset : token_offset + seq_len,
|
||||
slice_start_output:slice_end_output,
|
||||
].addmm_(x_slice, w_slice.T)
|
||||
|
||||
token_offset += seq_len
|
||||
|
||||
return output
|
||||
Reference in New Issue
Block a user