File size: 9,829 Bytes
2577656
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
# Adapted for diffusers from multimodal-art-projection/YuE at commit ef1936f2ee39fe8de486a0f47a481c95f8d4da87.
# Licensed under Apache-2.0; see LICENSE.
"""CUDA-graph token decoding for one request or its two guidance branches.

The graph replays the transformer's token stream for one new token per branch. The caller combines the branch logits,
samples once and passes that token to `step`. Each branch keeps its own RoPE positions and cache slots.
"""

from __future__ import annotations

import torch


def apply_rotary_emb(hidden_states, cos, sin):
    # Same arithmetic as `apply_rotary_emb` in the model repository's transformer.py.
    half = hidden_states.shape[-1] // 2
    x1, x2 = hidden_states[..., :half], hidden_states[..., half:]
    cos, sin = cos.to(hidden_states.dtype), sin.to(hidden_states.dtype)
    return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)


class _BranchPrefillCache:
    """Writes one branch's prompt keys and values into the graph's sequence-major buffers."""

    def __init__(self, keys, values, branch):
        self.keys = [buffer[branch : branch + 1] for buffer in keys]
        self.values = [buffer[branch : branch + 1] for buffer in values]
        self.seen_tokens = 0

    def get_seq_length(self):
        return self.seen_tokens

    def update(self, key, value, layer_idx):
        end = self.seen_tokens + key.shape[1]
        if end > self.keys[layer_idx].shape[1]:
            raise ValueError("Prefix exceeds preallocated KV capacity")
        self.keys[layer_idx][:, self.seen_tokens : end].copy_(key)
        self.values[layer_idx][:, self.seen_tokens : end].copy_(value)
        if layer_idx == len(self.keys) - 1:
            self.seen_tokens = end
        return self.keys[layer_idx][:, :end], self.values[layer_idx][:, :end]


class GraphAR:
    """Fixed-capacity decoding. `prefill()` returns the first logits; up to `max_tokens - 1` `step(token)` calls follow.

    Returned logits share output storage and stay valid until the next step.
    """

    def __init__(self, transformer, prefixes, max_tokens, device):
        if transformer.training:
            raise ValueError("GraphAR requires transformer.eval()")
        if not 1 <= len(prefixes) <= 2:
            raise ValueError("GraphAR supports one request or exactly two guidance branches")
        if device.type != "cuda":
            raise ValueError("CUDA graphs require a CUDA device")
        try:
            from diffusers.hooks.group_offloading import _get_group_onload_device

            _get_group_onload_device(transformer)
        except ValueError:
            pass
        else:
            raise ValueError(
                "CUDA graphs cannot capture a group-offloaded transformer: streamed weights change address between "
                "steps. Pass use_cuda_graph=False, or keep the transformer resident on the GPU."
            )
        config = transformer.config
        prefixes = [[int(token) for token in prefix] for prefix in prefixes]
        for prefix in prefixes:
            if not prefix or not all(0 <= token < config.vocab_size for token in prefix):
                raise ValueError("Prefixes require valid token IDs")
            if len(prefix) + max_tokens > config.max_position_embeddings:
                raise ValueError("Prefix plus generation budget exceeds model context; no length was shortened")
        self.transformer, self.device, self.dtype = transformer, device, transformer.dtype
        self.prefixes, self.max_tokens, self.branches = prefixes, int(max_tokens), len(prefixes)
        self.capacity = max(map(len, prefixes)) + self.max_tokens
        self.graph = self.output = None
        # PyTorch's variable-length FlashAttention takes per-branch key lengths on the GPU, so no mask is built.
        flash = (
            self.dtype in {torch.bfloat16, torch.float16}
            and config.attention_head_dim % 8 == 0
            and config.attention_head_dim <= 256
            and "seqused_k" in str(torch.ops.aten._flash_attention_forward.default._schema)
        )
        if not flash:
            raise ValueError(
                "CUDA graphs need PyTorch's variable-length FlashAttention with a BF16/FP16 transformer; "
                "pass use_cuda_graph=False."
            )
        self.ready, self.closed, self.steps = False, False, 0
        # Sequence-major buffers make the packed FlashAttention view contiguous without copying the cache each step.
        shape = (self.branches, self.capacity, config.num_key_value_heads, config.attention_head_dim)
        self.keys = [torch.zeros(shape, device=device, dtype=self.dtype) for _ in range(config.num_layers)]
        self.values = [torch.zeros(shape, device=device, dtype=self.dtype) for _ in range(config.num_layers)]
        self.positions = torch.tensor([len(prefix) for prefix in prefixes], dtype=torch.long, device=device)
        self.initial_positions = self.positions.clone()
        self.cu_q = torch.arange(self.branches + 1, dtype=torch.int32, device=device)
        self.cu_k = self.cu_q * self.capacity
        self.tokens = torch.tensor([[prefix[-1]] for prefix in prefixes], dtype=torch.long, device=device)

    @torch.inference_mode()
    def _decode(self):
        transformer = self.transformer
        config = transformer.config
        heads, kv_heads, head_dim = config.num_attention_heads, config.num_key_value_heads, config.attention_head_dim
        cos, sin = transformer.rotary_emb(self.positions[:, None])
        cos, sin = cos.unsqueeze(2), sin.unsqueeze(2)
        hidden_states = transformer.embed_tokens(self.tokens)
        used_lengths = (self.positions + 1).to(torch.int32)
        slots = self.positions[:, None, None, None].expand(self.branches, 1, kv_heads, head_dim)
        for block, keys, values in zip(transformer.transformer_blocks, self.keys, self.values):
            attn = block.attn
            normalized = block.norm1(hidden_states)
            query = attn.norm_q(attn.to_q(normalized).view(self.branches, 1, heads, head_dim))
            key = attn.norm_k(attn.to_k(normalized).view(self.branches, 1, kv_heads, head_dim))
            value = attn.to_v(normalized).view(self.branches, 1, kv_heads, head_dim)
            query = apply_rotary_emb(query, cos, sin)
            key = apply_rotary_emb(key, cos, sin)
            keys.scatter_(1, slots, key)
            values.scatter_(1, slots, value)
            # Every slot is allocated, but only each branch's filled prefix and current token are visible.
            # The packed variable-length entry point respects seqused_k; the 4D fixed-batch one ignores it.
            attn_output = torch.ops.aten._flash_attention_forward(
                query[:, 0],
                keys.view(-1, kv_heads, head_dim),
                values.view(-1, kv_heads, head_dim),
                self.cu_q,
                self.cu_k,
                1,
                self.capacity,
                0.0,
                False,
                False,
                seqused_k=used_lengths,
            )[0][:, None]
            hidden_states = hidden_states + attn.to_out[0](attn_output.reshape(self.branches, 1, -1))
            hidden_states = hidden_states + block.ff(block.norm2(hidden_states))
        output = transformer.lm_head(transformer.norm_out(hidden_states))[:, 0]
        self.positions.add_(1)
        return output

    @torch.inference_mode()
    def _capture(self):
        with torch.cuda.device(self.device):
            current = torch.cuda.current_stream(self.device)
            warmup = torch.cuda.Stream(device=self.device)
            warmup.wait_stream(current)
            with torch.cuda.stream(warmup):
                for _ in range(3):
                    self.positions.copy_(self.initial_positions)
                    self._decode()
                self.positions.copy_(self.initial_positions)
            current.wait_stream(warmup)
            torch.cuda.synchronize(self.device)
            self.graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(self.graph):
                self.output = self._decode()
            # Warmup and capture wrote one future slot; the first real step overwrites it before it becomes visible.
            self.positions.copy_(self.initial_positions)

    @torch.inference_mode()
    def prefill(self):
        if self.closed or self.ready:
            raise RuntimeError("prefill must be called exactly once on an open GraphAR")
        logits = []
        for branch, prefix in enumerate(self.prefixes):
            cache = _BranchPrefillCache(self.keys, self.values, branch)
            ids = torch.tensor([prefix], dtype=torch.long, device=self.device)
            logits.append(self.transformer(ids, kv_cache=cache, logits_to_keep=1).logits[:, -1])
        if any(param.device != self.device for param in self.transformer.parameters()):
            raise ValueError("CUDA graphs need the whole transformer resident on the GPU; pass use_cuda_graph=False.")
        if self.max_tokens > 1:
            self._capture()
        self.ready = True
        return torch.cat(logits, dim=0)

    @torch.inference_mode()
    def step(self, token):
        if self.closed or not self.ready:
            raise RuntimeError("Call prefill before step and do not use a closed GraphAR")
        if self.steps >= self.max_tokens - 1:
            raise ValueError("Requested generation budget is exhausted")
        self.tokens.copy_(token.reshape(1, 1).expand(self.branches, 1))
        self.graph.replay()
        self.steps += 1
        return self.output

    def close(self):
        self.graph = self.output = None
        self.keys.clear()
        self.values.clear()
        self.tokens = self.positions = self.initial_positions = None
        self.cu_q = self.cu_k = None
        self.closed = True