diff --git a/modules/attention/sparse/flex.py b/modules/attention/sparse/flex.py index 42ef1dad3..a9c1f979c 100644 --- a/modules/attention/sparse/flex.py +++ b/modules/attention/sparse/flex.py @@ -14,7 +14,9 @@ def to_block_mask(selection: BlockSelection, device=None) -> BlockMask: keep = keep.to(device) if keep.dim() != 4: raise ValueError(f'block selection must be 4d, got {tuple(keep.shape)}') - empty_num, empty_indices = _dense_to_ordered(torch.zeros_like(keep)) + # the partial slots stay empty by construction, so build them directly rather than sorting a mask of zeros + empty_num = torch.zeros(keep.shape[:-1], dtype=torch.int32, device=keep.device) + empty_indices = torch.zeros(keep.shape, dtype=torch.int32, device=keep.device) full_num, full_indices = _dense_to_ordered(keep) return BlockMask.from_kv_blocks( empty_num, empty_indices, diff --git a/modules/attention/sparse/selector.py b/modules/attention/sparse/selector.py index 89539e8cb..dc3cb19fc 100644 --- a/modules/attention/sparse/selector.py +++ b/modules/attention/sparse/selector.py @@ -24,12 +24,15 @@ class BlockSelection: budget: float seq_q: int seq_kv: int - density: float @property def shape(self) -> tuple[int, int, int, int]: return tuple(self.keep.shape) + def density(self) -> float: + """Fraction of tiles kept. Reads back from the accelerator, so this is for reporting and tests, never the hot path.""" + return float(self.keep.sum().item()) / max(self.keep.numel(), 1) + def block_count(length: int, block: int) -> int: return (length + block - 1) // block @@ -73,11 +76,15 @@ def score_blocks(query: torch.Tensor, key: torch.Tensor, spec: SparseSpec) -> to return torch.cat([pooled_q[:, i:i + chunk] @ pooled_k[:, i:i + chunk].transpose(-1, -2) for i in range(0, heads, chunk)], dim=1) -def select_blocks(query: torch.Tensor, key: torch.Tensor, spec: SparseSpec, pins: torch.Tensor | None = None, drops: torch.Tensor | None = None) -> BlockSelection | None: - """Keep the highest scoring KV tiles per query tile within the budget, plus pins and the diagonal. None means attend densely.""" - seq_q, seq_kv = query.shape[-2], key.shape[-2] - nq, nk = block_count(seq_q, spec.block_q), block_count(seq_kv, spec.block_kv) - device = query.device +plan_cache: dict = {} + + +def selection_plan(spec: SparseSpec, nq: int, nk: int, pins, drops, device, cache_key=None): + """The parts that depend only on geometry and layout, not on the tensors: what must be kept, what may be chosen, and how many.""" + key = (cache_key, nq, nk, spec.block_q, spec.block_kv, spec.budget, str(device)) + hit = plan_cache.get(key) if cache_key is not None else None + if hit is not None: + return hit must = diagonal_blocks(nq, nk, spec.block_q, spec.block_kv, device).unsqueeze(0).unsqueeze(0) if pins is not None: must = must | pins @@ -85,22 +92,33 @@ def select_blocks(query: torch.Tensor, key: torch.Tensor, spec: SparseSpec, pins candidates = ~must & ~forbidden per_row = candidates.sum(dim=-1, keepdim=True) # (.., NQ, 1) keep_per_row = torch.ceil(per_row * spec.budget).to(torch.int64) - if not spec.force and bool((keep_per_row >= per_row).all()): + covers_everything = bool((keep_per_row >= per_row).all()) # one readback, amortized over the generation by the cache + built = (must, forbidden, candidates, keep_per_row, covers_everything) + if cache_key is not None: + if len(plan_cache) > 32: + plan_cache.clear() + plan_cache[key] = built + return built + + +def select_blocks(query: torch.Tensor, key: torch.Tensor, spec: SparseSpec, pins: torch.Tensor | None = None, drops: torch.Tensor | None = None, cache_key=None) -> BlockSelection | None: + """Keep the highest scoring KV tiles per query tile within the budget, plus pins and the diagonal. None means attend densely.""" + seq_q, seq_kv = query.shape[-2], key.shape[-2] + nq, nk = block_count(seq_q, spec.block_q), block_count(seq_kv, spec.block_kv) + device = query.device + must, forbidden, candidates, keep_per_row, covers_everything = selection_plan(spec, nq, nk, pins, drops, device, cache_key) + if covers_everything and not spec.force: return None # the budget covers every candidate, so the mask would be dense scores = score_blocks(query, key, spec) scores = scores.masked_fill(~candidates.expand_as(scores), float('-inf')) - limit = int(keep_per_row.max().item()) - keep = must.expand(scores.shape).clone() - if limit > 0: - order = scores.argsort(dim=-1, descending=True, stable=True) - rank = torch.empty_like(order) - rank.scatter_(-1, order, torch.arange(nk, device=device).expand_as(order)) - keep |= (rank < keep_per_row) & candidates + # rank rather than topk, so the per row budget varies without a host side k + order = scores.argsort(dim=-1, descending=True, stable=True) + rank = torch.empty_like(order) + rank.scatter_(-1, order, torch.arange(nk, device=device).expand_as(order)) + keep = must | ((rank < keep_per_row) & candidates) keep &= ~forbidden - keep_int8 = keep.to(torch.int8) - density = float(keep_int8.sum().item()) / max(keep_int8.numel(), 1) - return BlockSelection(keep=keep_int8, block_q=spec.block_q, block_kv=spec.block_kv, budget=spec.budget, seq_q=seq_q, seq_kv=seq_kv, density=density) + return BlockSelection(keep=keep.to(torch.int8), block_q=spec.block_q, block_kv=spec.block_kv, budget=spec.budget, seq_q=seq_q, seq_kv=seq_kv) def radial_blocks(seq_q: int, seq_kv: int, density: float, spec: SparseSpec, device) -> BlockSelection: @@ -117,7 +135,7 @@ def radial_blocks(seq_q: int, seq_kv: int, density: float, spec: SparseSpec, dev else: high = mid keep = (distance <= high).unsqueeze(0).unsqueeze(0).to(torch.int8) - return BlockSelection(keep=keep, block_q=spec.block_q, block_kv=spec.block_kv, budget=density, seq_q=seq_q, seq_kv=seq_kv, density=float(keep.sum().item()) / max(keep.numel(), 1)) + return BlockSelection(keep=keep, block_q=spec.block_q, block_kv=spec.block_kv, budget=density, seq_q=seq_q, seq_kv=seq_kv) def schedule(steps: int, budget: float, bump: float = 0.0, bump_steps: int = 0) -> tuple[float, ...]: diff --git a/modules/attention/sparse/stage.py b/modules/attention/sparse/stage.py index 43d58256d..17922722f 100644 --- a/modules/attention/sparse/stage.py +++ b/modules/attention/sparse/stage.py @@ -103,7 +103,7 @@ def make_stage(options: StageOptions): if pins.shape[-2:] != (nq, nk): return decline('layout geometry mismatch') stage.last_skip = None - return select_blocks(query, key, spec, pins=pins, drops=drops) + return select_blocks(query, key, spec, pins=pins, drops=drops, cache_key=token_layout.key()) stage.options = options stage.last_skip = None diff --git a/test/test-attention-sparse.py b/test/test-attention-sparse.py index ca2ea5243..d78625da8 100644 --- a/test/test-attention-sparse.py +++ b/test/test-attention-sparse.py @@ -258,7 +258,7 @@ def test_radial_control_matches_the_requested_density(): spec = sparse.SparseSpec() for density in (0.15, 0.30): control = sparse.radial_blocks(4096, 4096, density, spec, device) - assert abs(control.density - density) < 0.05, f'requested {density}, got {control.density}' + assert abs(control.density() - density) < 0.05, f'requested {density}, got {control.density()}' return True @@ -410,7 +410,7 @@ def test_stage_uses_a_published_layout_and_falls_back_without_one(): loose = stage(q, k, v, None, False) ctx.set_layout(sparse.layout_from_segments([('text', 256), ('video', 1792)])) pinned = stage(q, k, v, None, False) - assert pinned.density > loose.density, f'pinning conditioning must keep more tiles: {pinned.density} vs {loose.density}' + assert pinned.density() > loose.density(), f'pinning conditioning must keep more tiles: {pinned.density()} vs {loose.density()}' assert bool(pinned.keep[..., 0:4].all()), 'the pinned text columns must survive' return True return with_context(checks)