"""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)