File size: 4,933 Bytes
f36843f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
"""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)