Skip to content

Commit 31ab3cd

Browse files
ChangyiYangnamezhenzhangclaude
committed
Phase A step 1: kill 2 syncs/forward in model LVM heads + EOS cache
Cherry-picked from upstream PR UCSB-AI#3 (namezhenzhang): * lvm_value_utils.get_eos_token_ids: cache result on req._lvm_eos_token_ids; was re-walking sets every candidate filter call. * qwen2_lvm.py / qwen3_lvm.py / qwen2_5_vl_lvm.py: drop two forward_batch.{extend_seq_lens,extend_prefix_lens}.tolist() calls per LVM forward by carrying tree_value_cached_prefix_lens in the TreeValueSpecInput. With ~1300 LVM forwards in the paper run, that removes ~2600 GPU->CPU syncs. * tree_value_spec.py: vectorize the candidate self-attention diagonal using numpy fancy indexing instead of a per-token Python loop. No algorithm change; same forward outputs. Co-Authored-By: Zhen Zhang <namezhenzhang@users.noreply.github.com> Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
1 parent e48cc37 commit 31ab3cd

5 files changed

Lines changed: 42 additions & 26 deletions

File tree

sglang-LenVM/python/sglang/srt/lvm/lvm_value_utils.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,10 @@ def get_eos_token_ids(req: Any) -> Set[int]:
1010
- `req.eos_token_ids`: a Set[int] from `ModelConfig.hf_eos_token_id` (preferred)
1111
- `req.tokenizer.eos_token_id`: tokenizer-defined EOS id
1212
"""
13+
cached = getattr(req, "_lvm_eos_token_ids", None)
14+
if cached is not None:
15+
return set(cached)
16+
1317
eos: Set[int] = set()
1418

1519
eos_ids = getattr(req, "eos_token_ids", None)
@@ -32,6 +36,10 @@ def get_eos_token_ids(req: Any) -> Set[int]:
3236
except Exception as exc:
3337
raise ValueError(f"Invalid LenVM tokenizer eos_token_id: {tok_eos!r}") from exc
3438

39+
try:
40+
setattr(req, "_lvm_eos_token_ids", tuple(sorted(eos)))
41+
except Exception:
42+
pass
3543
return eos
3644

3745

sglang-LenVM/python/sglang/srt/lvm/tree_value_spec.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -250,8 +250,8 @@ def build_tree_value_custom_mask_and_positions(
250250
# Candidate rows: attend all prefix tokens (0..L-1) and itself.
251251
cand_row_start = L - P
252252
m[cand_row_start : cand_row_start + N, :L] = True
253-
for j in range(N):
254-
m[cand_row_start + j, L + j] = True
253+
cand_arange = np.arange(N)
254+
m[cand_row_start + cand_arange, L + cand_arange] = True
255255

256256
mask_off += q_len * k_len
257257

sglang-LenVM/python/sglang/srt/models/qwen2_5_vl_lvm.py

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -77,22 +77,24 @@ def forward(
7777
)
7878

7979
if prefix_lens is not None and cand_lens is not None:
80-
extend_lens = forward_batch.extend_seq_lens.tolist()
80+
# Same upstream-PR-#3 sync removal as qwen2_lvm.py.
8181
cached_prefix_lens = (
82-
forward_batch.extend_prefix_lens.tolist()
83-
if forward_batch.extend_prefix_lens is not None
84-
else [0] * len(extend_lens)
82+
getattr(spec, "tree_value_cached_prefix_lens", None)
83+
or [0] * len(prefix_lens)
8584
)
8685

8786
out: list[torch.Tensor] = []
8887
offset = 0
89-
for i, ext_len in enumerate(extend_lens):
88+
for i, (prefix_len, cand_len, cached_prefix_len) in enumerate(
89+
zip(prefix_lens, cand_lens, cached_prefix_lens)
90+
):
91+
prefix_len = int(prefix_len)
92+
cand_len = int(cand_len)
93+
cached_prefix_len = int(cached_prefix_len)
94+
ext_len = max(prefix_len - cached_prefix_len, 0) + cand_len
9095
vals_i = token_values[offset : offset + ext_len]
9196
offset += ext_len
9297

93-
prefix_len = int(prefix_lens[i])
94-
cand_len = int(cand_lens[i])
95-
cached_prefix_len = int(cached_prefix_lens[i])
9698
cand_offset = max(prefix_len - cached_prefix_len, 0)
9799
out.append(vals_i[cand_offset : cand_offset + cand_len])
98100
return EmbeddingPoolerOutput(embeddings=out)

sglang-LenVM/python/sglang/srt/models/qwen2_lvm.py

Lines changed: 12 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -86,22 +86,26 @@ def forward(
8686
prefix_lens = getattr(spec, "tree_value_prefix_lens", None) if spec is not None else None
8787
cand_lens = getattr(spec, "tree_value_candidate_lens", None) if spec is not None else None
8888
if prefix_lens is not None and cand_lens is not None:
89-
extend_lens = forward_batch.extend_seq_lens.tolist()
89+
# Avoid forward_batch.{extend_seq_lens,extend_prefix_lens}.tolist()
90+
# syncs on every LVM forward by carrying these lengths in the
91+
# TreeValueSpecInput. Sourced from upstream PR #3.
9092
cached_prefix_lens = (
91-
forward_batch.extend_prefix_lens.tolist()
92-
if forward_batch.extend_prefix_lens is not None
93-
else [0] * len(extend_lens)
93+
getattr(spec, "tree_value_cached_prefix_lens", None)
94+
or [0] * len(prefix_lens)
9495
)
9596

9697
out: list[torch.Tensor] = []
9798
offset = 0
98-
for i, ext_len in enumerate(extend_lens):
99+
for i, (prefix_len, cand_len, cached_prefix_len) in enumerate(
100+
zip(prefix_lens, cand_lens, cached_prefix_lens)
101+
):
102+
L = int(prefix_len)
103+
N = int(cand_len)
104+
P = int(cached_prefix_len)
105+
ext_len = max(L - P, 0) + N
99106
vals_i = token_values[offset : offset + ext_len]
100107
offset += ext_len
101108

102-
L = int(prefix_lens[i])
103-
N = int(cand_lens[i])
104-
P = int(cached_prefix_lens[i])
105109
cand_offset = max(L - P, 0)
106110
out.append(vals_i[cand_offset : cand_offset + N])
107111
return EmbeddingPoolerOutput(embeddings=out)

sglang-LenVM/python/sglang/srt/models/qwen3_lvm.py

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -83,22 +83,24 @@ def forward(
8383
prefix_lens = getattr(spec, "tree_value_prefix_lens", None) if spec is not None else None
8484
cand_lens = getattr(spec, "tree_value_candidate_lens", None) if spec is not None else None
8585
if prefix_lens is not None and cand_lens is not None:
86-
extend_lens = forward_batch.extend_seq_lens.tolist()
86+
# Same upstream-PR-#3 sync removal as qwen2_lvm.py.
8787
cached_prefix_lens = (
88-
forward_batch.extend_prefix_lens.tolist()
89-
if forward_batch.extend_prefix_lens is not None
90-
else [0] * len(extend_lens)
88+
getattr(spec, "tree_value_cached_prefix_lens", None)
89+
or [0] * len(prefix_lens)
9190
)
9291

9392
out: list[torch.Tensor] = []
9493
offset = 0
95-
for i, ext_len in enumerate(extend_lens):
94+
for i, (prefix_len, cand_len, cached_prefix_len) in enumerate(
95+
zip(prefix_lens, cand_lens, cached_prefix_lens)
96+
):
97+
L = int(prefix_len)
98+
N = int(cand_len)
99+
P = int(cached_prefix_len)
100+
ext_len = max(L - P, 0) + N
96101
vals_i = token_values[offset : offset + ext_len]
97102
offset += ext_len
98103

99-
L = int(prefix_lens[i])
100-
N = int(cand_lens[i])
101-
P = int(cached_prefix_lens[i])
102104
cand_offset = max(L - P, 0)
103105
out.append(vals_i[cand_offset : cand_offset + N])
104106
return EmbeddingPoolerOutput(embeddings=out)

0 commit comments

Comments
 (0)