File size: 14,293 Bytes
9a25493
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
"""Local runtime with mutable, fixed-shape text conditioning for CUDA graphs.

For the released SEED model (text cross-RoPE disabled), each query sees exactly
one prompt segment. Repeated segments can therefore share that prompt's token
span. A bounded GPU token bank keeps forward shapes fixed while arbitrary
encoded prompts enter and leave the visible history. No model source changes.
"""
from __future__ import annotations

import torch

from space.inference import _signature
from window_runtime import WindowRuntime


def _feature_key(feature):
    try:
        version = feature._version
    except RuntimeError:  # Inference tensors have no version counter.
        version = None
    return id(feature), version, tuple(feature.shape), feature.dtype


class FixedPromptBank:
    """One-stream text-context adapter; all mutations use the caller's stream.

    Eviction follows the actual attention window, not the longer retained ring
    history. Slot moves only change masks and values, never the bank shape.
    RoPE-enabled or oversized contexts retain the official implementation and
    are explicitly marked for eager execution, without on-switch graph capture.
    """

    def __init__(self, text_module, original_context, buckets=(32, 64, 128, 256, 512, 1024)):
        buckets = tuple(buckets)
        if (not buckets or any(isinstance(n, bool) or not isinstance(n, int) or n <= 0 for n in buckets)
                or tuple(sorted(set(buckets))) != buckets):
            raise ValueError("Text token buckets must be increasing positive integers")
        self.text = text_module
        self.original_context = original_context
        self.buckets = buckets
        self.capacity = buckets[0]
        self.bank = None
        self.spans = {}
        self.keys = {}
        self._states = {}
        self.uploads = self.compactions = self.fallback_calls = 0
        self.max_active_tokens = self.active_tokens = 0
        self.last_mode = "uninitialized"
        self.null_length = int(self.text.text_cache[""].shape[0])

    @property
    def graph_compatible(self):
        return self.last_mode == "fixed_bank"

    def _fallback(self, reason, start, end, device, dtype):
        self.last_mode = reason
        self.fallback_calls += 1
        return self.original_context(start, end, device, dtype)

    def _first_free_span(self, length):
        cursor = 0
        for begin, size in sorted(self.spans.values()):
            if begin - cursor >= length:
                return cursor
            cursor = begin + size
        return cursor if self.capacity - cursor >= length else None

    def _upload(self, name, begin, length, device, dtype):
        feature = self.text.text_cache[name]
        self.bank[begin:begin + length].copy_(feature.to(device=device, dtype=dtype))
        self.spans[name] = (begin, length)
        self.keys[name] = _feature_key(feature)
        self.uploads += 1

    def __call__(self, start, end, device, dtype):
        if self.text.cross_rope:
            return self._fallback("cross_rope_eager", start, end, device, dtype)
        if len(self.text.stream_segments) != 1:
            return self._fallback("batch_eager", start, end, device, dtype)
        if int(self.text.text_cache[""].shape[0]) != self.null_length:
            return self._fallback("changed_null_length_eager", start, end, device, dtype)
        segments = self.text.stream_segments[0]
        ends = [index for _, index in segments[1:]] + [self.text.stream_frames[0]]
        active = [(name, begin, finish) for (name, begin), finish in zip(segments, ends)
                  if begin < end and finish > start]
        if not active:
            return self._fallback("empty_window_eager", start, end, device, dtype)
        names = list(dict.fromkeys(name for name, _, _ in active))
        lengths = {name: int(self.text.text_cache[name].shape[0]) for name in names}
        self.active_tokens = sum(lengths.values())
        self.max_active_tokens = max(self.max_active_tokens, self.active_tokens)
        capacity = next((size for size in self.buckets if size >= self.active_tokens), None)
        if capacity is None:
            return self._fallback("capacity_overflow_eager", start, end, device, dtype)

        width = int(self.text.text_cache[names[0]].shape[1])
        wanted_shape = (capacity, width)
        device = torch.device(device)
        if capacity not in self._states:
            self._states[capacity] = (torch.zeros(wanted_shape, device=device, dtype=dtype), {}, {})
        bank, spans, keys = self._states[capacity]
        if bank.shape != wanted_shape or bank.device != device or bank.dtype != dtype:
            return self._fallback("changed_tensor_contract_eager", start, end, device, dtype)
        self.capacity, self.bank, self.spans, self.keys = capacity, bank, spans, keys

        # Drop only unused or resized spans. Stale values are harmless because
        # every query's mask selects exactly its current prompt's valid span.
        for name in list(self.spans):
            if name not in lengths or self.spans[name][1] != lengths[name]:
                del self.spans[name]
                del self.keys[name]
        for name in names:
            length = lengths[name]
            key = _feature_key(self.text.text_cache[name])
            if name in self.spans:
                # Unversioned inference tensors are refreshed so an in-place
                # feature edit cannot leave captured conditioning stale.
                if key != self.keys[name] or key[1] is None:
                    self._upload(name, self.spans[name][0], length, device, dtype)
                continue
            begin = self._first_free_span(length)
            if begin is None:
                # Fragmentation, not overflow: stage a complete compact layout
                # from genuine cached features while preserving bank identity.
                self.spans.clear()
                self.keys.clear()
                cursor = 0
                for active_name in names:
                    self._upload(active_name, cursor, lengths[active_name], device, dtype)
                    cursor += lengths[active_name]
                self.compactions += 1
                break
            self._upload(name, begin, length, device, dtype)

        mask = torch.zeros(end - start, self.capacity, dtype=torch.bool, device=device)
        for name, begin, finish in active:
            offset, length = self.spans[name]
            mask[max(0, begin - start):min(end, finish) - start, offset:offset + length] = True
        self.last_mode = "fixed_bank"
        return [self.bank], {"cross_attn_mask": [mask]}

    def status(self):
        return dict(buckets=list(self.buckets), capacity=self.capacity, active_tokens=self.active_tokens,
                    max_active_tokens=self.max_active_tokens, resident_prompts=len(self.spans),
                    feature_uploads=self.uploads, compactions=self.compactions,
                    fallback_calls=self.fallback_calls, mode=self.last_mode)


def _pad_motion_inputs(args, kwargs, max_motion_frames):
    """Canonicalize a short 1x1 masked motion window for a prepared graph.

    The WAN already pads embedded states to seq_len. Raw zero padding changes
    only invalid states (e.g. the embedding bias); explicit self-attention
    masks prevent every valid query from reading those added keys. Valid
    temporal RoPE positions remain 0..N-1. This helper declines other contracts.
    """
    if max_motion_frames is None or len(args) != 4 or args[3] != max_motion_frames:
        return args, kwargs, None
    if any(kwargs.get(name) is not None for name in ("y", "rope_ids", "text_k_rope_ids", "text_q_rope_ids")):
        return args, kwargs, None
    values, times, contexts = args[:3]
    self_masks, cross_masks = kwargs.get("attn_mask"), kwargs.get("cross_attn_mask")
    sequences = (values, times, contexts, self_masks, cross_masks)
    if any(not isinstance(sequence, (list, tuple)) for sequence in sequences):
        return args, kwargs, None
    count = len(values)
    if not count or any(len(sequence) != count for sequence in sequences):
        return args, kwargs, None
    lengths = []
    for value, time, context, self_mask, cross_mask in zip(*sequences):
        if (value.ndim != 4 or value.shape[-2:] != (1, 1) or
                not 0 < value.shape[1] <= max_motion_frames):
            return args, kwargs, None
        length = value.shape[1]
        if (time.ndim != 1 or time.shape[0] != length or context.ndim != 2 or
                self_mask.shape != (length, length) or self_mask.dtype != torch.bool or
                cross_mask.shape != (length, context.shape[0]) or cross_mask.dtype != torch.bool):
            return args, kwargs, None
        lengths.append(length)
    if all(length == max_motion_frames for length in lengths):
        return args, kwargs, None

    padded_values, padded_times, padded_self, padded_cross = [], [], [], []
    for value, time, _, self_mask, cross_mask in zip(*sequences):
        length = value.shape[1]
        padded = value.new_zeros((value.shape[0], max_motion_frames, 1, 1))
        padded[:, :length].copy_(value)
        padded_values.append(padded)
        padded = time.new_zeros(max_motion_frames)
        padded[:length].copy_(time)
        padded_times.append(padded)
        padded = self_mask.new_zeros((max_motion_frames, max_motion_frames))
        padded[:length, :length].copy_(self_mask)
        padded_self.append(padded)
        padded = cross_mask.new_zeros((max_motion_frames, cross_mask.shape[1]))
        padded[:length].copy_(cross_mask)
        padded_cross.append(padded)
    changed_args = (padded_values, padded_times, contexts, args[3])
    changed_kwargs = dict(kwargs, attn_mask=padded_self, cross_attn_mask=padded_cross)
    return changed_args, changed_kwargs, lengths


class _BankGraphGate:
    def __init__(self, graph, bank, max_motion_frames=None):
        self.graph, self.bank = graph, bank
        self.max_motion_frames = max_motion_frames
        self.allow_capture = True
        self.uncached_eager_calls = 0
        self.startup_padded_calls = 0

    def __call__(self, *args, **kwargs):
        if self.bank.graph_compatible:
            if not self.graph.enabled:
                return self.graph(*args, **kwargs)
            graph_args, graph_kwargs, lengths = _pad_motion_inputs(args, kwargs, self.max_motion_frames)
            inputs = (graph_args, graph_kwargs)
            key = (_signature(inputs), torch.is_autocast_enabled("cuda"),
                   torch.get_autocast_dtype("cuda"))
            if self.allow_capture or key in self.graph.entries:
                output = self.graph(*graph_args, **graph_kwargs)
                if lengths is not None:
                    self.startup_padded_calls += 1
                    return [value[:, :length] for value, length in zip(output, lengths)]
                return output
            self.uncached_eager_calls += 1
        # A supported prompt is never truncated/rejected to fit the fast path.
        # Preserve existing graph entries for when a large history ages out.
        self.graph.eager_calls += 1
        return self.graph.forward(*args, **kwargs)


class FixedTextWindowRuntime(WindowRuntime):
    """WindowRuntime plus a dynamic prompt bank with fixed capacity buckets.

    Add genuine prompt features with the inherited add_prompt_features even
    during a session. New prompts fitting the active token budget reuse the
    same captured shapes. Bucket capacities are explicit and configurable at
    start. Overflow is correct eager inference, reported in status(). Prewarm
    all bucket graphs then freeze_graph_shapes() to prohibit mid-session capture.
    """

    def __init__(self, model, recovery_type, metadata):
        super().__init__(model, recovery_type, metadata)
        self._original_text_context = None
        self._fixed_text_bank = None
        self._bank_graph_gate = None

    @torch.inference_mode()
    def start(self, seed=0, history=120, use_graph=True, *, text_token_buckets=(32, 64, 128, 256, 512, 1024)):
        buckets = tuple(text_token_buckets)
        # Validate before resetting a running session.
        if (not buckets or any(isinstance(n, bool) or not isinstance(n, int) or n <= 0 for n in buckets)
                or tuple(sorted(set(buckets))) != buckets):
            raise ValueError("Text token buckets must be increasing positive integers")
        super().start(seed=seed, history=history, use_graph=use_graph)
        text = self.model.text_module
        self._original_text_context = text.get_stream_context
        self._fixed_text_bank = FixedPromptBank(text, self._original_text_context, buckets)
        text.get_stream_context = self._fixed_text_bank
        self._graph.max_entries = max(self._graph.max_entries, len(buckets))
        pad_history = (history if tuple(self.model.spatial_shape) == (1, 1) and
                       tuple(self.model.model.patch_size) == (1, 1, 1) else None)
        self._bank_graph_gate = _BankGraphGate(self._graph, self._fixed_text_bank, pad_history)
        self.model.model.forward = self._bank_graph_gate
        return self

    def freeze_graph_shapes(self):
        """Forbid new captures after prewarming; unknown shapes execute eagerly."""
        if not self._ready or self._bank_graph_gate is None:
            raise RuntimeError("Call start before freezing graph shapes")
        self._bank_graph_gate.allow_capture = False

    def close(self):
        if self._original_text_context is not None:
            self.model.text_module.get_stream_context = self._original_text_context
            self._original_text_context = None
        super().close()

    def status(self):
        status = super().status()
        status["fixed_text"] = self._fixed_text_bank.status() if self._fixed_text_bank else None
        if self._bank_graph_gate is not None:
            status["fixed_text"].update(capture_frozen=not self._bank_graph_gate.allow_capture,
                                      uncached_eager_calls=self._bank_graph_gate.uncached_eager_calls,
                                      startup_padded_calls=self._bank_graph_gate.startup_padded_calls)
        return status