perf: project only sampled rows through lm_head during prefill
- add logits_positions to AutoRegressiveLM.forward, gathering rows before the final norm so the lm_head GEMM covers only the positions prefill samples from - execute_prefill builds last_token_indices up front and passes them in, dropping the post-forward gather of a [tokens, vocab] tensor - prefill graph warmup passes a single index; decode stays untouched (every row is sampled) and prefill itself runs eager, so graph capture is unaffected - update the ragged-prefill fake to slice by the received index and add a packed-row exact-equality test Benchmark: NVIDIA L20 (idle), CUDA 12.8, torch 2.11.0+cu128, 1.2B bf16 checkpoint, 512-token prompts, greedy; prefill B=32: 368.1 -> 323.5 ms (44.5k -> 50.6k tok/s, +13.8%), B=8: 89.3 -> 78.9 ms (+13.2%), B=1: 12.3 -> 11.4 ms (+7.9%); decode step unchanged; full suite: 897 passed
This commit is contained in:
@@ -114,6 +114,53 @@ def test_model_forward_contract_uses_dense_training_and_packed_inference():
|
||||
)
|
||||
|
||||
|
||||
def test_forward_logits_positions_projects_only_requested_rows():
|
||||
"""logits_positions gathers packed rows before the lm_head projection."""
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.workspace import InferenceWorkspace
|
||||
|
||||
config = AutoRegressiveLMConfig(**TINY_CONFIG)
|
||||
model = AutoRegressiveLM(config).eval()
|
||||
prompts = [[1, 2, 3], [4, 5]]
|
||||
last_rows = torch.tensor([len(prompts[0]) - 1, len(prompts) - 1 + len(prompts[1])])
|
||||
|
||||
pool = PagePool(
|
||||
n_layers=config.num_hidden_layers,
|
||||
n_kv_heads=config.num_key_value_heads,
|
||||
head_dim=config.hidden_size // config.num_attention_heads,
|
||||
max_batch_size=2,
|
||||
max_seq_len=config.max_position_embeddings,
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
)
|
||||
cache = TaskCacheManager(pool)
|
||||
workspace = InferenceWorkspace(
|
||||
2,
|
||||
config.max_position_embeddings,
|
||||
config.num_attention_heads,
|
||||
config.hidden_size // config.num_attention_heads,
|
||||
torch.device("cpu"),
|
||||
torch.float32,
|
||||
)
|
||||
for tid, ids in zip(("t1", "t2"), prompts):
|
||||
assert cache.task_alloc(tid, ids)
|
||||
input_ids = torch.tensor(sum(prompts, []), dtype=torch.long)
|
||||
position_ids = torch.cat([torch.arange(len(p)) for p in prompts])
|
||||
with torch.inference_mode():
|
||||
kwargs = dict(
|
||||
position_ids=position_ids,
|
||||
kv_cache=cache.bind(["t1", "t2"], workspace, start_pos=0),
|
||||
fwd="prefill",
|
||||
)
|
||||
full = model(input_ids, **kwargs)
|
||||
sliced = model(input_ids, logits_positions=last_rows, **kwargs)
|
||||
|
||||
assert full["logits"].shape == (5, config.vocab_size)
|
||||
assert sliced["logits"].shape == (2, config.vocab_size)
|
||||
assert torch.equal(sliced["logits"], full["logits"][last_rows])
|
||||
assert torch.equal(sliced["hidden_states"], full["hidden_states"][last_rows])
|
||||
|
||||
|
||||
def _router_stats(probs, topk_indices):
|
||||
return {"probs": probs, "topk_indices": topk_indices}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user