File size: 8,608 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
197
198
199
200
201
# Adapted for diffusers from multimodal-art-projection/YuE at commit ef1936f2ee39fe8de486a0f47a481c95f8d4da87.
# Licensed under Apache-2.0; see LICENSE.
from __future__ import annotations

import time

import torch

from .protocol import ABC_END, CODEC_OFFSET, CODEC_SIZE, CONTEXT, EOD, MUSIC_END


class YuE2StaticKVCache:
    """Preallocated token cache for eager decoding; returns views of the filled prefix without copying history."""

    def __init__(self, num_layers, batch_size, num_kv_heads, max_seq_len, head_dim, dtype, device):
        self.num_layers = num_layers
        self.max_seq_len = max_seq_len
        self.seen_tokens = 0
        shape = (batch_size, num_kv_heads, max_seq_len, head_dim)
        self.keys = [torch.zeros(shape, dtype=dtype, device=device) for _ in range(num_layers)]
        self.values = [torch.zeros(shape, dtype=dtype, device=device) for _ in range(num_layers)]

    def get_seq_length(self):
        return self.seen_tokens

    def update(self, key, value, layer_idx):
        start, end = self.seen_tokens, self.seen_tokens + key.shape[1]
        if end > self.max_seq_len:
            raise ValueError(f"KV cache capacity {self.max_seq_len} exceeded by {end}; generation was not shortened")
        self.keys[layer_idx][:, :, start:end] = key.transpose(1, 2)
        self.values[layer_idx][:, :, start:end] = value.transpose(1, 2)
        if layer_idx == self.num_layers - 1:
            self.seen_tokens = end
        return self.keys[layer_idx][:, :, :end].transpose(1, 2), self.values[layer_idx][:, :, :end].transpose(1, 2)


def synchronize(device):
    if device.type == "cuda":
        torch.cuda.synchronize(device)
    elif device.type == "mps":
        torch.mps.synchronize()


def window_penalty(logits, recent_ids, penalty):
    if penalty == 1.0 or len(recent_ids) == 0:
        return logits
    recent = torch.as_tensor(recent_ids, dtype=torch.long, device=logits.device).reshape(1, -1)
    freq = torch.zeros_like(logits)
    freq.scatter_add_(-1, recent, torch.ones_like(recent, dtype=logits.dtype))
    alpha = penalty**freq
    return torch.where(logits < 0, logits * alpha, logits / alpha)


def distribution(logits, sampling, history, step, phase, legacy_off=False):
    # Planning-off requests keep the release's BF16 logits and top-3 floor; other modes sample from FP32 logits.
    scores = logits.clone() if legacy_off else logits.float().clone()
    end = ABC_END if phase == "abc" else MUSIC_END
    allowed = torch.full_like(scores, float("-inf"))
    if phase == "abc":
        allowed[..., :EOD] = 0
    else:
        allowed[..., CODEC_OFFSET : CODEC_OFFSET + CODEC_SIZE] = 0
    allowed[..., end] = 0
    scores = scores + allowed
    if step < sampling.min_tokens:
        scores[..., end] = -torch.inf
    scores = window_penalty(scores, history[-sampling.penalty_window :], sampling.repetition_penalty)
    if sampling.temperature == 0:
        return scores
    if sampling.temperature != 1:
        scores = scores / sampling.temperature
    threshold = scores.topk(min(sampling.top_k, scores.shape[-1])).values[..., -1, None]
    scores = scores.masked_fill(scores < threshold, -torch.inf)
    if sampling.top_p < 1:
        values, indices = scores.sort(descending=True)
        probabilities = values.softmax(-1)
        removed = probabilities.cumsum(-1) - probabilities > sampling.top_p
        removed[..., : 3 if legacy_off else 1] = False
        values = values.masked_fill(removed, -torch.inf)
        scores = values.scatter(-1, indices, values)
    return scores


@torch.inference_mode()
def generate_tokens(
    transformer,
    prefix,
    sampling,
    seed,
    phase,
    device,
    negative=None,
    combine_logits=None,
    legacy_off=False,
    cancelled=None,
    on_token=None,
    graph_decoder=None,
):
    """Sample one stage's tokens.

    With `negative`, `combine_logits(conditional, unconditional)` applies guidance. `graph_decoder` (the `GraphAR`
    class) decodes with CUDA graphs on CUDA devices; elsewhere decoding stays eager.
    """
    if len(prefix) + sampling.max_tokens > CONTEXT:
        raise ValueError("Prefix + requested generation budget exceeds 24576; no implicit truncation")
    if negative is not None and (combine_logits is None or len(negative) + sampling.max_tokens > CONTEXT):
        raise ValueError("Guidance needs `combine_logits` and a negative prefix that fits the context")
    if cancelled is not None and cancelled():
        raise InterruptedError("Cancelled before prefill")
    # Both stages reset the request seed, as the release does.
    rng_device = device if device.type in {"cpu", "cuda"} else torch.device("cpu")
    generator = torch.Generator(device=rng_device).manual_seed(seed)
    config = transformer.config

    def prefill(ids):
        cache = YuE2StaticKVCache(
            num_layers=config.num_layers,
            batch_size=1,
            num_kv_heads=config.num_key_value_heads,
            max_seq_len=len(ids) + sampling.max_tokens,
            head_dim=config.attention_head_dim,
            dtype=transformer.dtype,
            device=device,
        )
        logits = transformer(torch.tensor([ids], device=device), kv_cache=cache, logits_to_keep=1).logits
        return logits[:, -1, :], cache

    graph = None
    positive_cache = negative_cache = None
    synchronize(device)
    start = time.perf_counter()
    try:
        if graph_decoder is not None and device.type == "cuda":
            graph = graph_decoder(
                transformer, [prefix] if negative is None else [prefix, negative], sampling.max_tokens, device
            )
            logits = graph.prefill()
            conditional = logits[:1]
            unconditional = logits[1:] if negative is not None else None
        else:
            conditional, positive_cache = prefill(prefix)
            unconditional = None
            if negative is not None:
                unconditional, negative_cache = prefill(negative)
        synchronize(device)
        prefill_seconds = time.perf_counter() - start
        history, first, eos = [], None, False
        end = ABC_END if phase == "abc" else MUSIC_END
        for step in range(sampling.max_tokens):
            if cancelled is not None and cancelled():
                raise InterruptedError(f"Cancelled during {phase}")
            logits = conditional if negative is None else combine_logits(conditional, unconditional)
            scores = distribution(logits, sampling, history, step, phase, legacy_off)
            if sampling.temperature == 0:
                next_id = scores.argmax(-1, keepdim=True)
            else:
                probabilities = scores.softmax(-1)
                if device.type == "mps":
                    next_id = torch.multinomial(probabilities.cpu(), 1, generator=generator).to(device)
                else:
                    next_id = torch.multinomial(probabilities, 1, generator=generator)
            token = int(next_id.item())
            if first is None:
                first = time.perf_counter() - start
            if on_token is not None:
                on_token(phase, token)
            if token == end:
                eos = True
                break
            history.append(token)
            if step + 1 < sampling.max_tokens:
                if graph is not None:
                    branch_logits = graph.step(next_id)
                    conditional = branch_logits[:1]
                    unconditional = branch_logits[1:] if negative is not None else None
                else:
                    conditional = transformer(next_id, kv_cache=positive_cache, logits_to_keep=1).logits[:, -1, :]
                    if negative_cache is not None:
                        unconditional = transformer(next_id, kv_cache=negative_cache, logits_to_keep=1).logits[
                            :, -1, :
                        ]
        synchronize(device)
        seconds = time.perf_counter() - start
        count = len(history) + int(eos)
        timing = {
            "seconds": seconds,
            "prefill_seconds": prefill_seconds,
            "ttft_seconds": first,
            "output_tokens": count,
            "content_tokens": len(history),
            "output_tps": count / seconds,
            "prefix_tokens": len(prefix),
            "cfg_branches": 1 if negative is None else 2,
            "execution": "cuda_graph" if graph is not None else "eager",
        }
        return history, timing, not eos
    finally:
        if graph is not None:
            graph.close()
        positive_cache = negative_cache = None