Qwen-2.5-1B-RLCD-Fast / tree_decode.py
epsilon3's picture
Add tree-attention fast path and CUDA/M4 benchmarks
f36843f verified
Raw History Blame Contribute Delete
4.93 kB
"""Independent field branches sharing a single Qwen2 prefix KV cache.
Requires transformers 4.57.6; attention implementation must be sdpa or eager.
No cross-field conditioning: each field sees the prefix and its own ancestors.
"""
from dataclasses import dataclass
import torch
from transformers.cache_utils import DynamicCache
def fork_cache(prefix, fields=1):
# DynamicLayer.update concatenates, so shared prefix tensors are not mutated.
return DynamicCache(ddp_cache_data=[
(layer.keys if fields == 1 else layer.keys.repeat_interleave(fields, 0),
layer.values if fields == 1 else layer.values.repeat_interleave(fields, 0))
for layer in prefix.layers
])
@dataclass
class FieldPlan:
suffixes: list
prefix_length: int
device: str
dtype: torch.dtype
pad_id: int = 0
def __post_init__(self):
if not self.suffixes or any(not s for s in self.suffixes):
raise ValueError('Each field needs a nonempty suffix')
self.count = len(self.suffixes)
self.lengths = torch.tensor([len(s) for s in self.suffixes], device=self.device)
self.width = max(map(len, self.suffixes))
self.batch_ids = torch.tensor([
s + [self.pad_id] * (self.width - len(s)) for s in self.suffixes
], device=self.device)
self.batch_positions = (self.prefix_length + torch.arange(self.width, device=self.device))[None].expand(self.count, -1)
valid = torch.arange(self.width, device=self.device)[None] < self.lengths[:, None]
self.batch_mask = torch.cat((torch.ones(self.count, self.prefix_length, device=self.device, dtype=torch.bool), valid), 1)
self.tree_ids = torch.tensor([sum(self.suffixes, [])], device=self.device)
self.branches = torch.repeat_interleave(torch.arange(self.count, device=self.device), self.lengths)
self.depths = torch.cat([torch.arange(len(s), device=self.device) for s in self.suffixes])
self.tree_positions = (self.prefix_length + self.depths)[None]
self.leaves = self.lengths.cumsum(0) - 1
self.tree_mask = self.mask(self.branches, self.depths, self.branches, self.depths)
def mask(self, query_branches, query_depths, key_branches, key_depths):
allowed = (query_branches[:, None] == key_branches[None]) & (query_depths[:, None] >= key_depths[None])
allowed = torch.cat((torch.ones(len(query_branches), self.prefix_length, device=self.device, dtype=torch.bool), allowed), 1)
return torch.zeros(allowed.shape, device=self.device, dtype=self.dtype).masked_fill_(~allowed, torch.finfo(self.dtype).min)[None, None]
@torch.inference_mode()
def prefill(model, ids):
return model.model(ids, use_cache=True).past_key_values
@torch.inference_mode()
def decode_fields(model, prefix, plan, mode='tree', steps=1, forced_tokens=None):
"""Return [steps, fields, vocab] logits; optionally teacher-force continuation.
Both paths project only field endpoints through the same full LM head.
Timing includes cache branching, continuation masks, and token selection.
Fixed steps deliberately do not stop at EOS (benchmark workload).
"""
if mode not in ('tree', 'batch') or steps < 1:
raise ValueError('mode must be tree/batch and steps must be positive')
tree = mode == 'tree'
cache = fork_cache(prefix, 1 if tree else plan.count)
ids = plan.tree_ids if tree else plan.batch_ids
positions = plan.tree_positions if tree else plan.batch_positions
mask = plan.tree_mask if tree else plan.batch_mask
branches, depths = plan.branches, plan.depths
outputs = []
for step in range(steps):
hidden = model.model(input_ids=ids, position_ids=positions, attention_mask=mask,
past_key_values=cache, use_cache=True).last_hidden_state
if tree:
last = hidden[0, plan.leaves] if step == 0 else hidden[0]
else:
last = hidden[torch.arange(plan.count, device=plan.device), plan.lengths - 1] if step == 0 else hidden[:, 0]
logits = model.lm_head(last)
outputs.append(logits)
if step + 1 == steps:
break
tokens = forced_tokens[step] if forced_tokens is not None else logits.argmax(-1)
next_depths = plan.lengths + step
if tree:
ids = tokens[None]
query_branches = torch.arange(plan.count, device=plan.device)
branches = torch.cat((branches, query_branches))
depths = torch.cat((depths, next_depths))
mask = plan.mask(query_branches, next_depths, branches, depths)
positions = (plan.prefix_length + next_depths)[None]
else:
ids = tokens[:, None]
mask = torch.cat((mask, torch.ones(plan.count, 1, dtype=mask.dtype, device=plan.device)), 1)
positions = (plan.prefix_length + next_depths)[:, None]
return torch.stack(outputs)