Enable embedding lookup/lora_a logic for chunked backend (#17692)

Co-authored-by: Bruce Wu <mogicianwu@fb.com>
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
Co-authored-by: Ethan (Yusheng) Su <yushengsu.thu@gmail.com>
This commit is contained in:
Bruce Wu
2026-03-16 11:37:58 -07:00
committed by GitHub
parent 061ec582bf
commit 70a6fb53af
15 changed files with 930 additions and 46 deletions

View File

@@ -201,6 +201,67 @@ def reference_sgmv_shrink(
return output
def reference_embedding_lora_a_shrink(
input_ids: torch.Tensor,
weights: torch.Tensor,
weight_indices: torch.Tensor,
seq_lengths: torch.Tensor,
lora_ranks: torch.Tensor,
vocab_size: int,
) -> torch.Tensor:
"""
Simple sequence-level reference implementation of embedding LoRA A shrink operation.
Args:
input_ids: (total_seq_len,) - Token IDs
weights: (num_loras, max_rank, vocab_size) - LoRA A embedding weights
weight_indices: LoRA idx for each sequence
seq_lengths: Length of each sequence
lora_ranks: LoRA rank for each LoRA adapters
vocab_size: Base vocabulary size
Returns:
output: (total_seq_len, max_rank) - Embedded features
"""
if weights.numel() == 0:
total_tokens = input_ids.shape[0]
return torch.zeros(total_tokens, 0, dtype=weights.dtype, device=weights.device)
total_tokens = input_ids.shape[0]
_, max_rank, _ = weights.shape
output = torch.zeros(
total_tokens, max_rank, dtype=weights.dtype, device=weights.device
)
token_offset = 0
for lora_idx, seq_len, rank in zip(
weight_indices,
seq_lengths,
lora_ranks[weight_indices],
):
if seq_len == 0:
continue
if rank > 0:
# Get token IDs for this sequence
seq_input_ids = input_ids[token_offset : token_offset + seq_len]
# Clamp token IDs to vocab size
clamped_ids = torch.clamp(seq_input_ids, max=vocab_size - 1)
# Lookup embeddings: weights[lora_idx, :rank, token_ids] -> (seq_len, rank)
# weights shape: (num_loras, max_rank, vocab_size)
lora_weights = weights[lora_idx, :rank, :] # (rank, vocab_size)
embeddings = lora_weights[:, clamped_ids].t() # (seq_len, rank)
output[token_offset : token_offset + seq_len, :rank] = embeddings
token_offset += seq_len
return output
def reference_sgmv_expand(
x: torch.Tensor,
weights: torch.Tensor,