From 52c420f0b169e9bfeaba6a357cba5075a3557e96 Mon Sep 17 00:00:00 2001 From: Lingpeng Jin <103567126+valarLip@users.noreply.github.com> Date: Fri, 25 Sep 2026 03:39:55 +0000 Subject: [PATCH 1/5] perf: unify forward metadata publication and page-table reuse Publish persistent metadata at its first consumer boundary using checked direct copies or one packed H2D followed by scatter. Keep source lifetime, padding, stream identity, PP slots and explicit republication under one owner. Reject producer reentry before host writes and track actual packed-arena use. Share versioned page-table preparation across attention backends, including V4.1. Reuse unchanged maps, copy appended or reordered rows selectively, and record GPU revisions only after successful publication. Preserve worker append lineage without retaining page-id tuples or changing scheduler bookkeeping. Keep uniform sampling filters on the CPU, reuse independent TBO buffers, and combine Engram tentative cursors with snapshots. Group H2D regressions by contract and consumer; retain CPU page-map state tests separately. --- atom/model_engine/block_table_codec.py | 22 +- atom/model_engine/model_runner.py | 427 ++++++++--- atom/model_engine/sequence.py | 6 +- atom/model_ops/attentions/aiter_attention.py | 70 +- atom/model_ops/attentions/aiter_mla.py | 171 +++-- atom/model_ops/attentions/backends.py | 116 ++- .../attentions/deepseek_v41/backend.py | 83 ++- .../attentions/deepseek_v41/cache.py | 17 +- .../attentions/deepseek_v41/metadata.py | 86 ++- atom/model_ops/attentions/deepseek_v4_attn.py | 419 +++++++---- atom/model_ops/attentions/gdn_attn.py | 139 ++-- .../model_ops/attentions/kimi_mla_gdn_attn.py | 13 +- atom/model_ops/attentions/qwen4_exp_attn.py | 42 +- atom/model_ops/engram/device/hashing.py | 48 +- atom/model_ops/engram/device/runtime.py | 13 +- atom/model_ops/engram/device/staging.py | 10 +- atom/model_ops/sampler.py | 103 ++- atom/model_ops/v4_kernels/compress_plan.py | 55 +- atom/rollout/model_runner_ext.py | 4 +- atom/spec_decode/drafter.py | 89 ++- atom/utils/__init__.py | 25 +- atom/utils/block_tables.py | 233 ++++++ atom/utils/envs.py | 3 + atom/utils/forward_context.py | 8 +- atom/utils/h2d.py | 486 +++++++++++++ atom/utils/packed_h2d.py | 92 +++ docs/environment_variables.md | 6 + docs/h2d_publication.md | 134 ++++ tests/attentions/deepseek_v41/helpers.py | 48 ++ tests/attentions/deepseek_v41/test_cache.py | 47 +- .../attentions/deepseek_v41/test_metadata.py | 46 +- .../deepseek_v41/test_runtime_contract.py | 6 +- tests/model_ops/engram/test_overlap.py | 62 ++ .../deepseek_v41/test_dspark_integration.py | 13 +- .../deepseek_v41/test_sparse_attention.py | 9 +- .../deepseek_v41/test_speculative_state.py | 41 +- tests/test_h2d_attention_publication.py | 627 ++++++++++++++++ tests/test_h2d_draft_publication.py | 314 ++++++++ tests/test_h2d_publication.py | 678 ++++++++++++++++++ tests/test_h2d_runner_publication.py | 587 +++++++++++++++ tests/test_h2d_v4_indexer_publication.py | 634 ++++++++++++++++ tests/test_h2d_v4_publication.py | 419 +++++++++++ tests/test_model_runner_decode_padding.py | 171 +++++ tests/test_mtp_deferred_status_queue.py | 1 + tests/test_packed_h2d.py | 231 ++++++ tests/test_qwen4_exp_mtp.py | 24 +- tests/test_sampler_greedy_rows.py | 153 ++-- tests/test_sampler_scalar_filters.py | 86 +++ tests/test_shared_block_tables.py | 356 +++++++++ 49 files changed, 6676 insertions(+), 797 deletions(-) create mode 100644 atom/utils/block_tables.py create mode 100644 atom/utils/h2d.py create mode 100644 atom/utils/packed_h2d.py create mode 100644 docs/h2d_publication.md create mode 100644 tests/test_h2d_attention_publication.py create mode 100644 tests/test_h2d_draft_publication.py create mode 100644 tests/test_h2d_publication.py create mode 100644 tests/test_h2d_runner_publication.py create mode 100644 tests/test_h2d_v4_indexer_publication.py create mode 100644 tests/test_h2d_v4_publication.py create mode 100644 tests/test_model_runner_decode_padding.py create mode 100644 tests/test_packed_h2d.py create mode 100644 tests/test_sampler_scalar_filters.py create mode 100644 tests/test_shared_block_tables.py diff --git a/atom/model_engine/block_table_codec.py b/atom/model_engine/block_table_codec.py index 9b4f04df94..c78520883f 100644 --- a/atom/model_engine/block_table_codec.py +++ b/atom/model_engine/block_table_codec.py @@ -1,9 +1,9 @@ # SPDX-License-Identifier: MIT # Copyright (C) 2024-2025, Advanced Micro Devices, Inc. All rights reserved. -"""Ship a forward's block tables to the TP workers as appends alone. +"""Ship a forward's block tables to the workers as appends alone. -Every forward RPC broadcasts one `ScheduledBatch` to every TP worker, and its +Every forward RPC broadcasts one `ScheduledBatch` to every worker, and its `block_tables` are the bulk of it: one row per running request, the whole row every step, growing one block per decode. At 50 seqs x 100k context that is ~313k ids -- 1.2 MiB pickled and unpickled per rank per step -- to say @@ -25,7 +25,6 @@ as whole tables and resets both caches at once. """ -import array import copy import logging from dataclasses import dataclass @@ -158,10 +157,10 @@ def _send_whole(self, batch, reason: str): class BlockTableDeltaDecoder: - """Worker side: rebuild `array("i")` rows from a `BlockTableDelta`.""" + """Worker side: rebuild versioned rows from a `BlockTableDelta`.""" def __init__(self): - self._rows: dict[int, array.array] = {} + self._rows: dict[int, BlockTable] = {} def decode_rpc(self, func_name: str, args: list) -> list: """Decode `args[0]` in place if this is an encoded forward.""" @@ -184,8 +183,8 @@ def decode(self, batch): f"{len(req_ids)} requests" ) - rows: list[array.array] = [] - cached: dict[int, array.array] = {} + rows: list[BlockTable] = [] + cached: dict[int, BlockTable] = {} for i, req_id in enumerate(req_ids): req_id = int(req_id) base = int(delta.base_lengths[i]) @@ -203,9 +202,14 @@ def decode(self, batch): # rewrite history the token processor may still be reading, so # a row that grows is copied first. A row that did not grow is # shared, which is the decode-step-with-no-new-block case. - row = previous if end == start else previous[:] + row = previous + if end != start: + row = BlockTable(previous) + # This is the next immutable snapshot in the same append + # lineage. Equal version + length still means equal ids. + row.version = previous.version else: - row = array.array("i") + row = BlockTable() if end > start: row.frombytes(memoryview(delta.tail_values[start:end]).cast("B")) rows.append(row) diff --git a/atom/model_engine/model_runner.py b/atom/model_engine/model_runner.py index 5a7deb714d..1ce9ecb958 100644 --- a/atom/model_engine/model_runner.py +++ b/atom/model_engine/model_runner.py @@ -9,7 +9,7 @@ import os import time from contextlib import contextmanager, nullcontext -from functools import partial +from functools import partial, wraps from typing import Any, ClassVar, NamedTuple import numpy as np @@ -95,6 +95,7 @@ set_kv_cache_data, ) from atom.utils.gc_utils import freeze_gc_heap +from atom.utils.h2d import h2d_producer from atom.utils.selector import attn_family, get_attn_backend, has_mla_indexer from atom.utils.tbo import ( UBatchSlice, @@ -196,13 +197,19 @@ def __init__( self.runner = runner device = runner.device self.input_ids = CpuGpuBuffer( - max_num_batched_tokens + 1, dtype=torch.int32, device=device + max_num_batched_tokens + 1, + dtype=torch.int32, + device=device, + publication_group="input_ids", ) # One per request, not per token: where each request's anchor comes # from. Sized by tokens -- a batch can never hold more requests. The # matching prefix sum is `forward_vars["cu_seqlens_q"]`. self.decode_src = CpuGpuBuffer( - max_num_batched_tokens, dtype=torch.int32, device=device + max_num_batched_tokens, + dtype=torch.int32, + device=device, + publication_group="input_ids", ) self.use_spec = use_spec self.num_spec_tokens = num_spec_tokens @@ -456,10 +463,21 @@ def get_token_locations(self, batch: ScheduledBatch) -> TokenLocations: ), f"{n_deferred} deferred + {n_new} new != {num_cur} requests" return TokenLocations(deferred_curr, deferred_prev, new_curr) + def _publish_input_ids(self, count, group): + if group is None: + return self.input_ids.copy_to_gpu(count) + group.counts[group.indices["input_ids"]] = count + group.counts[group.indices["decode_src"]] = None + group.publish(group.counts) + return self.input_ids.gpu[:count] + + @h2d_producer("input_ids", runner="runner") def prepare_input_ids( self, batch: ScheduledBatch, max_seqlen_q: int, + *, + publication_group=None, ) -> torch.Tensor: """Prepare the input IDs for the current batch. @@ -472,12 +490,6 @@ def prepare_input_ids( total_tokens_prefill = batch.total_tokens_num_prefill total_tokens_decode = batch.total_tokens_num_decode total_reqs_prefill = batch.total_seqs_num_prefill - """for prefill: all input ids are new""" - self.input_ids.np[:total_tokens_prefill] = scheduled_tokens[ - :total_tokens_prefill - ] - self.input_ids.copy_to_gpu(total_tokens_prefill) - # The MTP status queue is filled in postprocess but drained here, so a # step whose postprocess is skipped must not drain it: `forward()` bails # before postprocess when the batch produces no output (every prefill in @@ -489,7 +501,12 @@ def prepare_input_ids( # TODO: remove this when we support mixed prefill and decode in one batch if total_reqs_prefill > 0: - return self.input_ids.gpu[:total_tokens_prefill] + # Decode does not publish an empty prefill prefix: an explicit + # zero copy is still a publication in this forward's ledger. + self.input_ids.np[:total_tokens_prefill] = scheduled_tokens[ + :total_tokens_prefill + ] + return self._publish_input_ids(total_tokens_prefill, publication_group) if not self.is_deferred_out: token_ids = scheduled_tokens[ @@ -501,7 +518,7 @@ def prepare_input_ids( raise NotImplementedError("pipeline parallel + speculative decode") self.input_ids.np[:total_tokens_decode] = token_ids - return self.input_ids.copy_to_gpu(total_tokens_decode) + return self._publish_input_ids(total_tokens_decode, publication_group) # PD consumer first decode: no prior prefill step initialized # prev_batch, so use scheduled_tokens directly for this step. @@ -510,7 +527,7 @@ def prepare_input_ids( total_tokens_prefill : total_tokens_prefill + total_tokens_decode ] self.input_ids.np[:total_tokens_decode] = token_ids - return self.input_ids.copy_to_gpu(total_tokens_decode) + return self._publish_input_ids(total_tokens_decode, publication_group) """for decode: input ids are from prev_sampled_token_ids""" locs = self.get_token_locations(batch) @@ -566,11 +583,18 @@ def prepare_input_ids( if n_draft > 0: s = int(cu_np[i]) + 1 self.input_ids.np[s : s + n_draft] = spec[i, :n_draft] - self.input_ids.copy_to_gpu(total_tokens_decode) - src_np = self.decode_src.np[:bs] src_np.fill(NEW_SEQUENCE) src_np[deferred_curr_indices] = deferred_prev_indices + group = ( + self.runner.h2d_groups["input_ids"] + if publication_group is None + else publication_group + ) + counts = group.counts + counts[group.indices["input_ids"]] = total_tokens_decode + counts[group.indices["decode_src"]] = bs + group.publish(counts) # How wide the forward reads: a replayed decode graph takes a fixed # `running_bs * tokens_per_seq` whatever the batch scheduled, and `bs` # sits between two captured buckets on most steps -- a 65-request batch @@ -585,7 +609,7 @@ def prepare_input_ids( fill_deferred_decode_ids( self.input_ids.gpu, self.runner.forward_vars["cu_seqlens_q"].gpu[: bs + 1], - self.decode_src.copy_to_gpu(bs), + self.decode_src.gpu[:bs], self.prev_token_ids, self.draft_token_ids if self.pre_num_decode_token_per_seq > 1 else None, max_tokens_per_seq=int(lens.max()) if bs else 1, @@ -799,11 +823,15 @@ def __init__(self, rank: int, config: Config): if getattr(self, "drafter", None) is not None: self.drafter.arm_aux_capture(self.model) self._init_forward_vars_ring() + self._init_h2d_publication() self.forward_done_event = torch.cuda.Event() initialize_eplb_runtime(self) self._maybe_warmup() - torch.set_default_device("cpu") + # Restore the implicit CPU default. An explicit "cpu" leaves a + # DeviceContext intercepting every torch call during serving, including + # tensor views and Triton's pointer/stride specialization queries. + torch.set_default_device(None) torch.set_default_dtype(default_dtype) if self.config.compilation_config.level == 1: @@ -1254,10 +1282,19 @@ def allocate_forward_vars(self): # TODO: remove it in forward_context self.forward_vars = { "input_ids": self.tokenID_processor.input_ids, - "positions": CpuGpuBuffer(self.max_num_batched_tokens, **i64_kwargs), - "temperatures": CpuGpuBuffer(self.max_bs, **f32_kwargs), - "top_ks": CpuGpuBuffer(self.max_bs, **i32_kwargs), - "top_ps": CpuGpuBuffer(self.max_bs, **f32_kwargs), + "decode_src": self.tokenID_processor.decode_src, + "positions": CpuGpuBuffer( + self.max_num_batched_tokens, publication_group="positions", **i64_kwargs + ), + "temperatures": CpuGpuBuffer( + self.max_bs, publication_group="sampling", **f32_kwargs + ), + "top_ks": CpuGpuBuffer( + self.max_bs, publication_group="sampling", **i32_kwargs + ), + "top_ps": CpuGpuBuffer( + self.max_bs, publication_group="sampling", **f32_kwargs + ), # Keep enough space for MTP decode (max_q_len > 1). # `extra_output_dims` lets a model insert dims between N and dim # (e.g. DeepSeek-V4 returns the un-reduced mHC residual @@ -1272,16 +1309,21 @@ def allocate_forward_vars(self): } if self.use_mrope: self.forward_vars["mrope_positions"] = CpuGpuBuffer( - 3, self.max_num_batched_tokens, **i64_kwargs + 3, + self.max_num_batched_tokens, + publication_group="mrope", + publication_unit="elements", + **i64_kwargs, ) if hasattr(self, "drafter"): self.forward_vars["mtp_k"] = self.drafter.mtp_k + self.forward_vars.update(self.drafter.metadata_buffers) self.forward_vars["num_accepted_tokens"] = CpuGpuBuffer( self.max_bs, **i32_kwargs ) # Per in-flight slot via forward_vars; PP ring clones it. self.forward_vars["draft_next_tokens"] = CpuGpuBuffer( - self.max_bs, **i32_kwargs + self.max_bs, publication_group="draft_anchors", **i32_kwargs ) def _init_forward_vars_ring(self): @@ -1338,6 +1380,69 @@ def _clone_slot(src: dict) -> dict: self._fv_slot_events = [torch.cuda.Event() for _ in range(pp_size)] logger.info(f"forward_vars ring: {pp_size} slots (pipeline parallel)") + def _init_h2d_publication(self): + from atom.utils.h2d import PublicationOwner, PublicationRegistry + + registry = PublicationRegistry() + self.publication_registry = registry + transport = envs.ATOM_H2D_BACKEND + if transport not in ("direct", "packed"): + raise ValueError("ATOM_H2D_BACKEND must be direct or packed") + self._h2d_owners = [] + self._h2d_groups = [] + for index, variables in enumerate(self._fv_ring): + event = ( + self._stage_h2d_done + if self._fv_slot_events is None + else self._fv_slot_events[index] + ) + owner = PublicationOwner(self.device, event, registry=registry) + members = {} + for name, buffer in variables.items(): + if isinstance(buffer, CpuGpuBuffer) and buffer.publication_group: + binding = owner.bind(buffer, name, unit=buffer.publication_unit) + members.setdefault(buffer.publication_group, []).append(binding) + groups = {name: owner.group(name, items) for name, items in members.items()} + if transport == "packed": + # One upload immediately before token assembly's first GPU + # consumer. Sampling has no earlier device consumer. + if all(name in groups for name in ("sampling", "early", "input_ids")): + token_members = ("sampling", "early", "input_ids") + if "spec_decode" in groups: + token_members += ("spec_decode",) + groups["token_inputs"] = owner.group( + "token_inputs", + [b for name in token_members for b in groups[name].members], + ) + if "prefill" in groups and "positions" in groups: + groups["prefill_inputs"] = owner.group( + "prefill_inputs", + groups["prefill"].members + groups["positions"].members, + ) + # Some buffers also publish together at a later consumer boundary + # (e.g. MHA decode shares block tables with the prefill group). + for name, names in getattr( + getattr(self, "attn_metadata_builder", None), "h2d_group_members", {} + ).items(): + groups[name] = owner.group( + name, [variables[item]._publication for item in names] + ) + if transport == "packed": + for group in owner.use_packed_transport(): + logger.info( + "H2D group %s: %s%s", + group.name, + group.transport, + f" ({group.fallback_reason})" if group.fallback_reason else "", + ) + # Cover constructor uploads, ring clones and transport pointer + # tables before the first preparation, even on another stream. + event.record() + self._h2d_owners.append(owner) + self._h2d_groups.append(groups) + self.h2d_owner = self._h2d_owners[self._fv_idx] + self.h2d_groups = self._h2d_groups[self._fv_idx] + def _advance_forward_vars(self): """Rotate to the next in-flight slot before any buffer is written. @@ -1347,14 +1452,15 @@ def _advance_forward_vars(self): if len(self._fv_ring) == 1: return self._fv_idx = (self._fv_idx + 1) % len(self._fv_ring) - # Block until this slot's previous forward finished reading it on the - # GPU before we overwrite its host-pinned staging buffers. No-op unless - # the CPU has raced > ring-size forwards ahead of the GPU. - self._fv_slot_events[self._fv_idx].synchronize() + # Select the slot now; _gate_staging_reuse waits on its existing event + # before any producer writes. Avoid waiting twice on the same event. self.forward_vars = self._fv_ring[self._fv_idx] - # `input_ids` is the one forward_vars buffer aliased outside the dict - # (tokenID_processor writes into it directly); repoint it at this slot. + self.h2d_owner = self._h2d_owners[self._fv_idx] + self.h2d_groups = self._h2d_groups[self._fv_idx] + # Token assembly holds aliases outside forward_vars; select both + # payload and source-index mirrors from this slot before writing. self.tokenID_processor.input_ids = self.forward_vars["input_ids"] + self.tokenID_processor.decode_src = self.forward_vars["decode_src"] def _gate_staging_reuse(self): """Block until the previous forward's staging H2Ds have executed. @@ -1380,33 +1486,33 @@ def _gate_staging_reuse(self): their copies are in flight. They must therefore enter this gate and record the event just like real forwards. - The pipeline ring solves the same problem by rotating buffers, which - bounds the lead to its depth; `_stage_h2d_done` is None there and this - does nothing. + The publication owner uses the existing single-slot or PP slot event, + then opens one ledger epoch before any producer writes host metadata. """ - if self._stage_h2d_done is not None: - self._stage_h2d_done.synchronize() + self.h2d_owner.begin() def _mark_staging_h2d_enqueued(self): """Close the window the gate above waits on. - Every `_stage` / `copy_to_gpu` a forward does is enqueued inside - `prepare_model` -- `build()` fences the current stream behind - `prep_stream` before returning -- so one event after it covers them - all. `prepare_mtp_decode` is the exception, staging from inside - `postprocess`, a path that synchronizes on its own. + Covers this preparation phase's source reads, including registered + direct copies and grouped publishers. Late V4 TBO preparation resumes + this epoch and records the completion again after its uploads. + Independent subsystem uploads retain their own completion protocols. """ - if self._stage_h2d_done is not None: - self._stage_h2d_done.record() + self.h2d_owner.finish() def _record_forward_vars_event(self): """Mark the current slot's forward as done on the GPU stream. Paired - with the synchronize() in ``_advance_forward_vars``. Called at the end of + with the owner wait in ``_gate_staging_reuse``. Called at the end of every forward, including DP-sync dummies. No-op when the ring has a single slot.""" if len(self._fv_ring) == 1: return - self._fv_slot_events[self._fv_idx].record() + try: + self._fv_slot_events[self._fv_idx].record() + except BaseException: + self.h2d_owner.fail() + raise def _get_num_kv_heads(self): """Return the per-rank number of KV heads.""" @@ -1596,7 +1702,7 @@ def get_num_blocks(self) -> dict[str, object]: # This prevents OOM when other processes share the GPU. available_for_kv = min(available_for_kv_budget, free) - torch.set_default_device("cpu") + torch.set_default_device(None) specs = self._sub_pool_specs() @@ -2347,6 +2453,8 @@ def prepare_inputs( batch: ScheduledBatch, input_ids: torch.Tensor, forward_mode: ForwardMode, + *, + spec_decode_indices: tuple[np.ndarray, int] | None = None, ): # Always supplied, settled in `prepare_model` (which is where the reason # lives). The q-bucket shrink ran there too, so `batch` is already @@ -2369,6 +2477,14 @@ def prepare_inputs( # sizes everything per-sequence, `running_tokens` everything per-row. running_bs = forward_mode.running_bs running_tokens = forward_mode.running_tokens + spec_decode_metadata = None + if not is_prefill and hasattr(self, "drafter") and not batch.is_dummy_run: + _, lens, cu = self.attn_metadata_builder.decode_spans(batch) + # Gather after token assembly, before attention/Engram's final + # consumers. Packed indices already share the token input upload. + spec_decode_metadata = self.drafter.calc_spec_decode_metadata( + lens, cu[1:], input_ids, prepared_indices=spec_decode_indices + ) attn_metadata, positions = self.attn_metadata_builder.build( batch=batch, running_bs=running_bs, @@ -2388,15 +2504,6 @@ def prepare_inputs( forward_mode=forward_mode, ) - spec_decode_metadata = None - if not is_prefill and hasattr(self, "drafter") and not batch.is_dummy_run: - _, lens, cu = self.attn_metadata_builder.decode_spans(batch) - # `cu[1:]` is the segment ENDS, which is what - # `cu_num_sampled_tokens` means. - spec_decode_metadata = self.drafter.calc_spec_decode_metadata( - lens, cu[1:], input_ids - ) - pcp_size = self.config.prefill_context_parallel_size _pcp_tbo_balanced = ( is_prefill @@ -2433,9 +2540,12 @@ def prepare_inputs( ub_tokens_across_dp=ub_tokens_across_dp, ) + @h2d_producer("sampling") def prepare_sample( - self, batch: ScheduledBatch - ) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None, bool, bool]: + self, batch: ScheduledBatch, *, publication_group=None + ) -> tuple[ + torch.Tensor, int | torch.Tensor | None, float | torch.Tensor | None, bool, bool + ]: bs = batch.total_seqs_num # Check on CPU whether all requests are greedy (temperature=0) @@ -2443,41 +2553,52 @@ def prepare_sample( # Check on CPU whether any fan-out sibling needs per-row random noise. # Missing attribute (e.g. dummy runs, older callers) -> False. - needs_independent_noise = bool( - getattr(batch, "needs_independent_noise", np.zeros(0, dtype=bool)).any() - ) + noise = getattr(batch, "needs_independent_noise", None) + needs_independent_noise = noise is not None and bool(noise.any()) temp_buffer = self.forward_vars["temperatures"] # Clamp temperatures on CPU to avoid division by zero in sampler - temp_buffer.np[:bs] = np.maximum(batch.temperatures, SAMPLER_EPS) - temperatures = temp_buffer.copy_to_gpu(bs) + np.maximum(batch.temperatures, SAMPLER_EPS, out=temp_buffer.np[:bs]) # Check on CPU whether filtering is needed to avoid GPU sync in sampler. # If no filtering needed, return None to skip GPU copy entirely. needs_topk = (batch.top_ks != -1).any() needs_topp = (batch.top_ps < 1.0).any() + # Uniform filters are already known on the CPU. Keep them as scalars + # for AITER's fast dispatch instead of uploading and reading them back. + top_ks = top_ps = None + top_k_count = top_p_count = None if needs_topk: - top_k_buffer = self.forward_vars["top_ks"] - top_k_buffer.np[:bs] = batch.top_ks - # If all values are the same, only copy one element to save bandwidth - if bs > 1 and (batch.top_ks == batch.top_ks[0]).all(): - top_ks = top_k_buffer.copy_to_gpu(1) + if bs == 1 or (batch.top_ks == batch.top_ks[0]).all(): + top_ks = int(batch.top_ks[0]) else: - top_ks = top_k_buffer.copy_to_gpu(bs) - else: - top_ks = None + top_k_buffer = self.forward_vars["top_ks"] + top_k_buffer.np[:bs] = batch.top_ks + top_ks = top_k_buffer.gpu[:bs] + top_k_count = bs if needs_topp: - top_p_buffer = self.forward_vars["top_ps"] - top_p_buffer.np[:bs] = batch.top_ps - # If all values are the same, only copy one element to save bandwidth - if bs > 1 and (batch.top_ps == batch.top_ps[0]).all(): - top_ps = top_p_buffer.copy_to_gpu(1) + if bs == 1 or (batch.top_ps == batch.top_ps[0]).all(): + top_ps = float(np.float32(batch.top_ps[0])) else: - top_ps = top_p_buffer.copy_to_gpu(bs) - else: - top_ps = None + top_p_buffer = self.forward_vars["top_ps"] + top_p_buffer.np[:bs] = batch.top_ps + top_ps = top_p_buffer.gpu[:bs] + top_p_count = bs + + group = ( + self.h2d_groups["sampling"] + if publication_group is None + else publication_group + ) + counts = group.counts + counts[group.indices["temperatures"]] = bs + counts[group.indices["top_ks"]] = top_k_count + counts[group.indices["top_ps"]] = top_p_count + if publication_group is None: + group.publish(counts) + temperatures = temp_buffer.gpu[:bs] return temperatures, top_ks, top_ps, all_greedy, needs_independent_noise @@ -2514,20 +2635,48 @@ def prepare_model(self, batch: ScheduledBatch): total_tokens_num = batch.total_tokens_num assert total_tokens_num > 0 + token_group = self.h2d_groups.get("token_inputs") + spec_decode_indices = None temperatures, top_ks, top_ps, all_greedy, needs_independent_noise = ( - self.prepare_sample(batch) + self.prepare_sample(batch, publication_group=token_group) ) - # Publishes the buffer `prepare_input_ids` addresses spans through. - self.attn_metadata_builder.publish_cu_seqlens_q(batch, forward_mode) - input_ids = self.tokenID_processor.prepare_input_ids( - batch, forward_mode.max_seqlen_q + if token_group is None: + self.attn_metadata_builder.publish_cu_seqlens_q(batch, forward_mode) + input_ids = self.tokenID_processor.prepare_input_ids( + batch, forward_mode.max_seqlen_q + ) + else: + cu_count = self.attn_metadata_builder.prepare_cu_seqlens_q( + batch, forward_mode + ) + token_group.counts[token_group.indices["cu_seqlens_q"]] = cu_count + spec_group = self.h2d_groups.get("spec_decode") + if spec_group is not None: + if batch.total_tokens_num_prefill == 0 and not batch.is_dummy_run: + _, lens, cu = self.attn_metadata_builder.decode_spans(batch) + spec_decode_indices = self.drafter.prepare_spec_decode_indices( + lens, cu[1:], token_group + ) + else: + # A preceding decode may have filled these counts. Prefill + # and dummy forwards must not publish its stale indices. + for binding in spec_group.members: + token_group.counts[token_group.indices[binding.name]] = None + # No GPU consumer runs between these host producers and this + # upload. Token assembly publishes before launching its kernel. + input_ids = self.tokenID_processor.prepare_input_ids( + batch, forward_mode.max_seqlen_q, publication_group=token_group + ) + self.prepare_inputs( + batch, + input_ids, + forward_mode=forward_mode, + spec_decode_indices=spec_decode_indices, ) - self.prepare_inputs(batch, input_ids, forward_mode=forward_mode) - # Stage the speculative inputs while this forward's normal staging - # window is still open. Both buffers are pinned and reused, so copying - # them later from postprocess would fall outside the event recorded by - # forward() immediately after prepare_model(). + # Stage scheduler anchor overrides while this forward's staging window + # is still open. This pinned mirror is reused, so copying it later from + # postprocess would fall outside the preparation completion event. if hasattr(self, "drafter"): forward_context = get_forward_context() if batch.next_token_ids is not None: @@ -2756,6 +2905,30 @@ def _is_shared(layer_idx): self._pp_index_topk, ) + def _padded_decode_inputs(self, forward_mode: ForwardMode): + """Expose the full decode run, independently of graph replay. + + Metadata and uniform DP collectives already use running_tokens. Keep + model activations at that height too, with legal ids/positions in the + unused tail. Mixed prefill/decode uses varlen collectives and must not + call this helper. + """ + assert not forward_mode.is_prefill + assert forward_mode.running_tokens_are_unified + scheduled = forward_mode.scheduled_tokens + running = forward_mode.running_tokens + ids = self.forward_vars["input_ids"].gpu[:running] + positions = ( + self._mrope_positions_view(running) + if self.use_mrope + else self.forward_vars["positions"].gpu[:running] + ) + assert ids.shape[0] == running and positions.shape[-1] == running + if running > scheduled: + ids[scheduled:].zero_() + positions[..., scheduled:].zero_() + return ids, positions + @record_gpu_forward def run_model( self, @@ -2818,6 +2991,8 @@ def run_model( # prefill, or decode forced eager (enforce_eager / DP peer # prefill / bs above the largest captured graph). with record_function(label): + if not is_prefill and forward_mode.running_tokens_are_unified: + input_ids, positions = self._padded_decode_inputs(forward_mode) # The multimodal runtime owns request leases and span scatter. inputs_embeds = None if ( @@ -2920,7 +3095,13 @@ def run_model( model_output = self._restore_pcp_balanced_output( model_output, _pcp_bal_groups, _pcp_size ) - hidden_states = model_output + # PP carries the full run between stages. Only the last + # stage returns scheduled rows to sampling/draft consumers. + hidden_states = ( + model_output[: context.scheduled_tokens] + if not is_prefill + else model_output + ) logits = self.model.compute_logits(hidden_states) else: # decode[bs=128 tok=128 d=128] / decode[... p=2 d=126 spec=3] / @@ -2931,23 +3112,7 @@ def run_model( scheduled_tokens = context.scheduled_tokens if self._piecewise_cg_active(): - # Pad tail to a legal vocab id / position, from THIS rank's - # own rows out to the width the step settled on. A group-max - # lower bound leaves `[scheduled, max)` holding the previous - # step's ids on every rank below the max, and those reach the - # draft's Markov lookup as out-of-range indices. - if running_tokens > scheduled_tokens: - self.forward_vars["input_ids"].gpu[ - scheduled_tokens:running_tokens - ].zero_() - self.forward_vars["positions"].gpu[ - scheduled_tokens:running_tokens - ].zero_() - _pos = ( - self._mrope_positions_view(running_tokens) - if self.use_mrope - else self.forward_vars["positions"].gpu[:running_tokens] - ) + _ids, _pos = self._padded_decode_inputs(forward_mode) forward_context.cudagraph_runtime_mode = ( CUDAGraphMode.PIECEWISE if forward_mode.piecewise_captured @@ -2956,9 +3121,7 @@ def run_model( forward_context.batch_descriptor = BatchDescriptor( num_tokens=running_tokens ) - model_output = self.model( - self.forward_vars["input_ids"].gpu[:running_tokens], _pos - ) + model_output = self.model(_ids, _pos) forward_context.cudagraph_runtime_mode = CUDAGraphMode.NONE forward_context.batch_descriptor = None # model_output is always a plain Tensor; drafter aux capture @@ -3000,8 +3163,8 @@ def postprocess( batch: ScheduledBatch, logits: torch.Tensor, temperatures: torch.Tensor, - top_ks: torch.Tensor | None, - top_ps: torch.Tensor | None, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, all_greedy: bool, # following for draft hidden_states: torch.Tensor, @@ -3842,7 +4005,9 @@ def pause_gc(): full_q_len, ) + self.h2d_owner.begin() self.attn_metadata_builder.blank_cache_write_targets() + self.h2d_owner.finish() # Whether this backend's capture builder supports a dynamic (per-bucket) build_capture = self.attn_metadata_builder.build_for_cudagraph_capture @@ -3851,6 +4016,34 @@ def pause_gc(): # Whether it supports a ragged num_tokens_pad (zero-copy-q attn-core graphs). supports_ragged_capture = "num_tokens_pad" in _build_params + raw_build_capture = build_capture + + @wraps(raw_build_capture) + def build_capture(*args, **kwargs): + # Preparation is outside actual graph capture, including calls + # made by the drafter and ragged bucket builder. Each invocation + # owns a synthetic transaction; normal forwards never use it. + bs = kwargs.get("bs", args[0] if args else None) + q_len = kwargs.get("max_q_len", full_q_len) + self.h2d_owner.begin() + try: + if not self.attn_metadata_builder.capture_owns_cu_seqlens_q: + cu = self.forward_vars["cu_seqlens_q"] + cu.np[: bs + 1] = np.arange( + 0, (bs + 1) * q_len, q_len, dtype=np.int32 + ) + cu.copy_to_gpu(bs + 1) + tokens = bs * q_len + self.forward_vars["positions"].np[:tokens] = ( + np.arange(tokens, dtype=np.int64) % q_len + ) + result = raw_build_capture(*args, **kwargs) + self.h2d_owner.finish() + return result + except BaseException: + self.h2d_owner.fail() + raise + with pause_gc(), graph_capture() as capture_ctx, self.capture_profiler as prof: for max_q_len in q_buckets: capture_range = ( @@ -3862,12 +4055,6 @@ def pause_gc(): if self.rank == 0: capture_range.set_description(f"Capturing {bs=}, {max_q_len=}") - cu_seqlens_q = np.arange( - 0, (bs + 1) * max_q_len, max_q_len, dtype=np.int32 - ) - self.forward_vars["cu_seqlens_q"].np[: bs + 1] = cu_seqlens_q - self.forward_vars["cu_seqlens_q"].copy_to_gpu(bs + 1) - num_tokens = bs * max_q_len if _piecewise and self._piecewise_skip_capture(num_tokens): continue @@ -3876,10 +4063,6 @@ def pause_gc(): # its handful of Python statements just fold into the next # iteration's window. self._capture_trace_tag = f"bs_{bs}_q_{max_q_len}" - # Use a simple, safe position pattern for capture. - self.forward_vars["positions"].np[:num_tokens] = ( - np.arange(num_tokens, dtype=np.int64) % max_q_len - ) if supports_dynamic_q_len: attn_metadata, context = build_capture( bs=bs, max_q_len=max_q_len @@ -4282,6 +4465,8 @@ def forward(self, batch: ScheduledBatch) -> ScheduledBatchOutput: self._done_event.record() stream.wait_event(self._done_event) with torch.cuda.stream(stream): + self._advance_forward_vars() + self._gate_staging_reuse() ( input_ids, temperatures, @@ -4290,6 +4475,7 @@ def forward(self, batch: ScheduledBatch) -> ScheduledBatchOutput: all_greedy, needs_independent_noise, ) = self.prepare_model(batch) + self._mark_staging_h2d_enqueued() logits, hidden_states = self.run_model(input_ids, batch) self._model_fwd_event.record(stream) torch.cuda.current_stream().wait_event(self._model_fwd_event) @@ -4306,6 +4492,7 @@ def forward(self, batch: ScheduledBatch) -> ScheduledBatchOutput: needs_independent_noise=needs_independent_noise, ) + self._record_forward_vars_event() reset_forward_context() return fwd_output @@ -4635,6 +4822,8 @@ def prefill_forward(self, batch: ScheduledBatch) -> list[int]: else: stream = torch.cuda.current_stream() with torch.cuda.stream(stream): + self._advance_forward_vars() + self._gate_staging_reuse() ( input_ids, temperatures, @@ -4643,10 +4832,12 @@ def prefill_forward(self, batch: ScheduledBatch) -> list[int]: all_greedy, _needs_independent_noise, ) = self.prepare_model(batch) + self._mark_staging_h2d_enqueued() logits, _ = self.run_model(input_ids, batch) # Sample the first generated token from each sequence's last logit sampled = self.sampler(logits, temperatures, top_ks, top_ps, all_greedy) sampled_cpu = sampled.view(-1).tolist() + self._record_forward_vars_event() # Synchronize so decode's default stream sees all KV writes. stream.synchronize() self._record_kv_cache_ready(batch) diff --git a/atom/model_engine/sequence.py b/atom/model_engine/sequence.py index a6368a1c9e..400e33634b 100644 --- a/atom/model_engine/sequence.py +++ b/atom/model_engine/sequence.py @@ -61,8 +61,10 @@ class BlockTable(array.array): method lookup and nothing else. A version is a global draw rather than a per-table counter so that a fresh - table for a recycled request id can never look like a known one: no two - live tables ever carry the same version. + table for a recycled request id can never look like a known one: + independent tables never share a version. The RPC decoder may retain + a version when it copies a row to append: those immutable snapshots share + one prefix lineage, and equal lengths still identify identical contents. """ __slots__ = ("version",) diff --git a/atom/model_ops/attentions/aiter_attention.py b/atom/model_ops/attentions/aiter_attention.py index 6c44b576c2..04f29691e9 100644 --- a/atom/model_ops/attentions/aiter_attention.py +++ b/atom/model_ops/attentions/aiter_attention.py @@ -19,6 +19,7 @@ block_table_convert_triton, kv_indices_generate_triton, ) +from atom.utils.block_tables import block_table_state from atom.utils.forward_context import AttentionMetaData, Context, get_forward_context from atom.utils.tbo import TokenSplitPrefillState @@ -324,13 +325,26 @@ def __init__( dtype=reduce_partial_map_type, device=self.device, ), - "kv_indptr": CpuGpuBuffer(self.max_bs + 1, **i32_kwargs), + "kv_indptr": CpuGpuBuffer( + self.max_bs + 1, publication_group="mha_csr", **i32_kwargs + ), "kv_indices": CpuGpuBuffer( self.max_bs * self.max_num_blocks_per_seq, **i32_kwargs, ), } self.model_runner.forward_vars.update(pa_persistent_metadata) + # Ready together before the CSR kernel, sharing prefill destinations + # without changing their addresses or padded publication counts. + self.h2d_group_members = { + "mha_decode": ( + "slot_mapping", + "context_lens", + "block_tables", + "kv_indptr", + "positions", + ), + } # Per-ubatch buffers for CUDAGraph TBO if model_runner.config.enable_tbo: self._allocate_ubatch_buffers( @@ -375,20 +389,22 @@ def _allocate_ubatch_buffers( for ub_idx in range(self._NUM_TBO_UBATCHES): p = f"ub{ub_idx}_" - var[f"{p}kv_indptr"] = CpuGpuBuffer(ub_max_bs + 1, **i32_kwargs) + host_i32 = dict(i32_kwargs, publication_group=f"{p}mha_metadata") + host_i64 = dict(i64_kwargs, publication_group=f"{p}mha_metadata") + var[f"{p}kv_indptr"] = CpuGpuBuffer(ub_max_bs + 1, **host_i32) var[f"{p}kv_indices"] = CpuGpuBuffer( self.max_bs * self.max_num_blocks_per_seq, **i32_kwargs, ) - var[f"{p}context_lens"] = CpuGpuBuffer(ub_max_bs, **i32_kwargs) + var[f"{p}context_lens"] = CpuGpuBuffer(ub_max_bs, **host_i32) var[f"{p}slot_mapping"] = CpuGpuBuffer( ub_max_bs * max_seqlen_qo, - **i64_kwargs, + **host_i64, ) var[f"{p}block_tables"] = CpuGpuBuffer( - ub_max_bs, self.block_table_cols, **i32_kwargs + ub_max_bs, self.block_table_cols, **host_i32 ) - var[f"{p}cu_seqlens_q"] = CpuGpuBuffer(ub_max_bs + 1, **i32_kwargs) + var[f"{p}cu_seqlens_q"] = CpuGpuBuffer(ub_max_bs + 1, **host_i32) var[f"{p}cu_seqlens_q"].cpu.copy_( torch.arange( 0, @@ -968,9 +984,9 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): ) if self._has_sparse_attention and not attn_metadata.has_cached: bs = batch.total_seqs_num_prefill - attn_metadata.block_tables = self.model_runner.forward_vars[ - "block_tables" - ].copy_to_gpu(bs) + attn_metadata.block_tables = block_table_state( + self.model_runner.forward_vars["block_tables"] + ).publish(bs) # `prefill_attention_triton` reads the paged KV cache, so it needs a # block_table even with no cached tokens. The base builder marshals one # every step but only uploads it when `has_cached`. @@ -980,9 +996,9 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): and batch.block_tables ): bs = batch.total_seqs_num_prefill - attn_metadata.block_tables = self.model_runner.forward_vars[ - "block_tables" - ].copy_to_gpu(bs) + attn_metadata.block_tables = block_table_state( + self.model_runner.forward_vars["block_tables"] + ).publish(bs) if self._has_sparse_attention: from atom.model_ops.minimax_m3.sparse_attn import ( make_sparse_prefill_metadata, @@ -1163,6 +1179,7 @@ def prepare_decode( running_tokens: int, max_seqlen_q: int, ): + self._check_metadata_writable("mha_csr", "prefill", "positions", "mrope") scheduled_bs = batch.total_seqs_num_decode self.total_blocks = 0 dropout_p = 0.0 @@ -1179,7 +1196,7 @@ def prepare_decode( max_seqlen_k = np.max(context_lens) # Before the slots, not after: `slot_mapping` reads this packed table. - self.prepare_block_tables(batch) + self.prepare_block_tables(batch, running_bs) var = self.model_runner.forward_vars scheduled_tokens = batch.total_tokens_num_decode @@ -1224,7 +1241,14 @@ def prepare_decode( ("kv_indptr", running_bs + 1), ] - ctx = {el: var[el].copy_to_gpu(num) for el, num in vars_used} + group = self.model_runner.h2d_groups["mha_decode"] + for name, count in vars_used: + group.counts[group.indices[name]] = count + group.counts[group.indices["positions"]] = ( + None if self.model_runner.use_mrope else scheduled_tokens + ) + block_table_state(var["block_tables"]).publish(running_bs, group=group) + ctx = {el: var[el].gpu[:num] for el, num in vars_used} # A view: `publish_cu_seqlens_q` already uploaded it this step, and # nothing here writes the host copy. ctx["cu_seqlens_q"] = var["cu_seqlens_q"].gpu[: running_bs + 1] @@ -1274,12 +1298,12 @@ def prepare_decode( n_valid_column_per_row_out=self._n_valid_column_per_row_buffer(), ) mrope_positions = self._build_mrope_decode_positions( - batch, context_lens, max_seqlen_q + batch, context_lens, max_seqlen_q, running_tokens=running_tokens ) if mrope_positions is not None: positions = mrope_positions else: - positions = var["positions"].copy_to_gpu(scheduled_tokens) + positions = var["positions"].gpu[:scheduled_tokens] if self.model_runner.config.enable_tbo_decode and running_bs >= 2: self._prepare_ubatch_decode( scheduled_bs, @@ -1302,6 +1326,7 @@ def _prepare_ubatch_decode( Splits the full-batch data into per-ubatch CpuGpuBuffers. The split point is bs // 2 to match CUDAGraph's baked-in token slices. """ + self._check_metadata_writable("ub0_mha_metadata", "ub1_mha_metadata") var = self.model_runner.forward_vars N = self._NUM_TBO_UBATCHES half = bs // N @@ -1328,10 +1353,9 @@ def _prepare_ubatch_decode( ] var[f"{p}slot_mapping"].np[ub_real_tokens:ub_running_tokens] = -1 - var[f"{p}block_tables"].np[:ub_real_reqs] = var["block_tables"].np[ - req_start : req_start + ub_real_reqs - ] - var[f"{p}block_tables"].np[ub_real_reqs:running_bs] = 0 + block_table_state(var["block_tables"]).slice_to( + var[f"{p}block_tables"], req_start, ub_real_reqs, pad_to=running_bs + ) full_kv_indptr = var["kv_indptr"].np base = full_kv_indptr[req_start] @@ -1361,8 +1385,10 @@ def _prepare_ubatch_decode( (f"{p}kv_indptr", running_bs + 1), (f"{p}cu_seqlens_q", running_bs + 1), ] - for el, num in vars_used: - var[el].copy_to_gpu(num) + group = self.model_runner.h2d_groups[f"{p}mha_metadata"] + for name, count in vars_used: + group.counts[group.indices[name]] = count + block_table_state(var[f"{p}block_tables"]).publish(running_bs, group=group) ub_max_seqlen_k = ( int(context_lens[req_start : req_start + ub_real_reqs].max()) diff --git a/atom/model_ops/attentions/aiter_mla.py b/atom/model_ops/attentions/aiter_mla.py index 48ccd740ee..94656f0ac2 100644 --- a/atom/model_ops/attentions/aiter_mla.py +++ b/atom/model_ops/attentions/aiter_mla.py @@ -65,6 +65,7 @@ kv_indices_generate_triton, mtp_prepare_decode_mla_kernel, ) +from atom.utils.block_tables import block_table_state from atom.utils.forward_context import AttentionMetaData, Context from .backends import AttentionBackend, CommonAttentionBuilder @@ -140,7 +141,7 @@ def aligned_index_cache_dim(hf_config) -> int: def _pad_prefill_mla_draft_tail( kv_indptr: torch.Tensor, kv_last_page_lens: np.ndarray, - block_tables: np.ndarray, + block_tables: np.ndarray | None, scheduled_bs: int, running_bs: int, ) -> None: @@ -150,7 +151,8 @@ def _pad_prefill_mla_draft_tail( return kv_indptr[scheduled_bs + 1 : running_bs + 1] = kv_indptr[scheduled_bs] kv_last_page_lens[scheduled_bs:running_bs] = 0 - block_tables[scheduled_bs:running_bs] = 0 + if block_tables is not None: + block_tables[scheduled_bs:running_bs] = 0 def _global_index_cache_layer_ids( @@ -515,36 +517,48 @@ def __init__(self, model_runner): dtype=reduce_partial_map_type, device=self.device, ), - "kv_indptr": CpuGpuBuffer(self.max_bs + 1, **i32_kwargs), + "kv_indptr": CpuGpuBuffer( + self.max_bs + 1, publication_group="mla_csr", **i32_kwargs + ), # Global (un-sharded) per-request KV indptr for round-robin CP: cumsum # of the GLOBAL context_lens (token-level, page_size=1). Only filled # when dcp_world_size > 1; consumed by the cprr kernel via # mla_decode_fwd(g_kv_indptr=...) to apply the global-position causal # mask for MTP (max_q_len > 1). - "g_kv_indptr": CpuGpuBuffer(self.max_bs + 1, **i32_kwargs), + "g_kv_indptr": CpuGpuBuffer( + self.max_bs + 1, publication_group="mla_csr", **i32_kwargs + ), "kv_indices": CpuGpuBuffer( self.max_bs * self.max_num_blocks_per_seq, **i32_kwargs, ), - "kv_last_page_lens": CpuGpuBuffer(self.max_bs, **i32_kwargs), + "kv_last_page_lens": CpuGpuBuffer( + self.max_bs, publication_group="mla_csr", **i32_kwargs + ), } if self._publishes_dcp_local_lens: # Layer-invariant sparse-DSA indexer metadata: one row per query # token, derived once per step and reused by every full layer. mla_metadata["dcp_local_context_lens"] = CpuGpuBuffer( - self.max_bs * max_seqlen_qo, **i32_kwargs + self.max_bs * max_seqlen_qo, publication_group="mla_csr", **i32_kwargs ) mla_metadata["kv_last_page_lens"].cpu.fill_(1) mla_metadata["kv_last_page_lens"].copy_to_gpu() if self.is_sparse: mla_metadata["cu_seqlen_ke"] = CpuGpuBuffer( - self.max_num_batched_tokens, **i32_kwargs + self.max_num_batched_tokens, + publication_group="mla_sparse", + **i32_kwargs, ) mla_metadata["cu_seqlen_ks"] = CpuGpuBuffer( - self.max_num_batched_tokens, **i32_kwargs + self.max_num_batched_tokens, + publication_group="mla_sparse", + **i32_kwargs, ) mla_metadata["sparse_kv_indptr"] = CpuGpuBuffer( - self.max_num_batched_tokens + 1, **i32_kwargs + self.max_num_batched_tokens + 1, + publication_group="mla_sparse", + **i32_kwargs, ) mla_metadata["sparse_cu_seqlens_q"] = CpuGpuBuffer( self.max_num_batched_tokens + 1, **i32_kwargs @@ -693,6 +707,10 @@ def __init__(self, model_runner): ) self.model_runner.forward_vars.update(mla_metadata) + if self.is_sparse: + self.model_runner.forward_vars["batch_id_per_q_token"].publication_group = ( + "mla_tokens" + ) # Chunked-context workspaces for the prefill has_cached path. Sized # to config.attn_prefill_chunk_size (defaults to max_num_batched_tokens) @@ -797,26 +815,28 @@ def _allocate_ubatch_buffers( for ub_idx in range(self._NUM_TBO_UBATCHES): p = f"ub{ub_idx}_" - var[f"{p}kv_indptr"] = CpuGpuBuffer(ub_max_bs + 1, **i32_kwargs) + host_i32 = dict(i32_kwargs, publication_group=f"{p}mla_metadata") + host_i64 = dict(i64_kwargs, publication_group=f"{p}mla_metadata") + var[f"{p}kv_indptr"] = CpuGpuBuffer(ub_max_bs + 1, **host_i32) # Per-ubatch global (un-sharded) kv_indptr for round-robin CP (see the # shared "g_kv_indptr" buffer). Filled in _build_ubatch when dcp>1. - var[f"{p}g_kv_indptr"] = CpuGpuBuffer(ub_max_bs + 1, **i32_kwargs) + var[f"{p}g_kv_indptr"] = CpuGpuBuffer(ub_max_bs + 1, **host_i32) var[f"{p}kv_indices"] = CpuGpuBuffer( self.max_bs * self.max_num_blocks_per_seq, **i32_kwargs, ) - var[f"{p}context_lens"] = CpuGpuBuffer(ub_max_bs, **i32_kwargs) - var[f"{p}kv_last_page_lens"] = CpuGpuBuffer(ub_max_bs, **i32_kwargs) + var[f"{p}context_lens"] = CpuGpuBuffer(ub_max_bs, **host_i32) + var[f"{p}kv_last_page_lens"] = CpuGpuBuffer(ub_max_bs, **host_i32) var[f"{p}kv_last_page_lens"].cpu.fill_(0) var[f"{p}kv_last_page_lens"].copy_to_gpu() var[f"{p}slot_mapping"] = CpuGpuBuffer( ub_max_bs * max_seqlen_qo, - **i64_kwargs, + **host_i64, ) var[f"{p}block_tables"] = CpuGpuBuffer( - ub_max_bs, self.block_table_cols, **i32_kwargs + ub_max_bs, self.block_table_cols, **host_i32 ) - var[f"{p}cu_seqlens_q"] = CpuGpuBuffer(ub_max_bs + 1, **i32_kwargs) + var[f"{p}cu_seqlens_q"] = CpuGpuBuffer(ub_max_bs + 1, **host_i32) var[f"{p}cu_seqlens_q"].cpu.copy_( torch.arange( 0, @@ -830,7 +850,7 @@ def _allocate_ubatch_buffers( if self.is_sparse: var[f"{p}sparse_kv_indptr"] = CpuGpuBuffer( ub_max_bs + 1, - **i32_kwargs, + **host_i32, ) # MLA work buffers per ubatch (GPU only) @@ -1774,7 +1794,18 @@ def _sparse_selected_counts(self, seq_lens): history = np.minimum(pools, self.index_topk // kpool) * kpool return (history + seq_lens % kpool).astype(np.int32) + def _prefill_block_table_rows(self, scheduled_bs, running_bs, has_cached): + if hasattr(self.model_runner, "drafter") or has_cached: + # Prepare the final padded range before the first upload. Later + # consumers reuse it, without rewriting a borrowed host tail. + block_table_state(self.model_runner.forward_vars["block_tables"]).pad( + scheduled_bs, running_bs + ) + return running_bs + return None + def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): + self._check_metadata_writable("mla_csr", "mla_sparse") attn_metadata, positions = CommonAttentionBuilder.prepare_prefill( self, batch, running_bs ) @@ -1790,13 +1821,17 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): and self.index_kpool > 1 and attn_metadata.block_tables is None ): - self.prepare_block_tables(batch) - attn_metadata.block_tables = var["block_tables"].copy_to_gpu(bs) + # Already packed by the common producer; publish its final layout. + attn_metadata.block_tables = block_table_state(var["block_tables"]).publish( + bs + ) if self.is_sparse and attn_metadata.max_seqlen_k > self.index_topk: if attn_metadata.block_tables is None: # Already marshalled by the base builder; only the upload is # gated on `has_cached`. - attn_metadata.block_tables = var["block_tables"].copy_to_gpu(bs) + attn_metadata.block_tables = block_table_state( + var["block_tables"] + ).publish(bs) counts = var["cu_seqlens_q"].np[1 : bs + 1] - var["cu_seqlens_q"].np[:bs] local_offsets = np.concatenate( [np.arange(s, dtype=np.int32) for s in counts] @@ -1933,19 +1968,18 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): ) # kv_indices_generate_triton expects logical block_tables (one - # entry per block_ratio tokens). The parent packed exactly that - # this step, and the only write to the mirror in between is the - # tail zeroing below, which starts at `bs`. + # entry per block_ratio tokens). The parent published the final + # padded table already; its host mirror remains borrowed. _pad_prefill_mla_draft_tail( kv_indptr, var["kv_last_page_lens"].np, - var["block_tables"].np, + None, # block-table padding was prepared before its only upload bs, running_bs, ) var["kv_last_page_lens"].copy_to_gpu(running_bs) attn_metadata.kv_last_page_lens = var["kv_last_page_lens"].gpu[:bs] - block_tables_for_kv = var["block_tables"].copy_to_gpu(running_bs)[:bs] + block_tables_for_kv = var["block_tables"].gpu[:bs] kv_indices_generate_triton( block_tables_for_kv, attn_metadata.kv_indices, @@ -2438,6 +2472,7 @@ def prepare_decode( running_tokens: int, max_seqlen_q: int, ): + self._check_metadata_writable("mla_csr", "mla_sparse", "prefill", "positions") scheduled_bs = batch.total_seqs_num_decode dropout_p = 0.0 @@ -2456,7 +2491,7 @@ def prepare_decode( # Before the slots, not after: `slot_mapping` reads this packed table. # DCP still walks the trimmed ragged rows -- its slot is a per-rank # filter, not an address this table can answer. - self.prepare_block_tables(batch) + self.prepare_block_tables(batch, running_bs) if not batch.is_dummy_run: if max_seqlen_q > 1: @@ -2602,7 +2637,7 @@ def prepare_decode( scheduled_tokens + 1 : running_tokens + 1 ] = var["sparse_kv_indptr"].np[scheduled_tokens] vars_used.append(("sparse_kv_indptr", running_tokens + 1)) - vars_used.append(("sparse_cu_seqlens_q", running_tokens + 1)) + # Immutable unit-stride prefix, uploaded once at allocation. metadata_deps.add("sparse_kv_indptr") else: sparse_context_lens = self._sparse_selected_counts( @@ -2618,44 +2653,27 @@ def prepare_decode( metadata_deps.add("sparse_kv_indptr") vars_for_metadata = [(el, num) for el, num in vars_used if el in metadata_deps] - vars_remaining = [(el, num) for el, num in vars_used if el not in metadata_deps] + vars_remaining = [ + (el, num) + for el, num in vars_used + if el not in metadata_deps and el != "block_tables" + ] max_seqlen_k = context_lens.max() - # The side prep_stream overlaps the metadata H2D copies + kv_indices - # generation with the main stream. Under intra-GPU disagg the decode runs - # on a CU-masked stream, and the prep_stream's wait_stream barriers - # serialize against it, adding per-step decode latency. So in disagg mode - # do the copies synchronously on the current stream; otherwise keep the - # async overlap. - disagg = self.model_runner.config.enable_rapidserve - ctx = {} - ctx["kv_indptr"] = var["kv_indptr"].copy_to_gpu(running_bs + 1) - if disagg: - ctx_rest = {el: var[el].copy_to_gpu(num) for el, num in vars_remaining} - ctx.update(ctx_rest) - ctx["kv_indices"] = var["kv_indices"].gpu - kv_indices_generate_triton( - ctx["block_tables"], - ctx["kv_indices"], - ctx["kv_indptr"], - self.block_ratio, - max_seqlen_k, - ) - else: - prep_stream = self.prep_stream - current_stream = torch.cuda.current_stream() - prep_stream.wait_stream(current_stream) - with torch.cuda.stream(prep_stream): - ctx_rest = {el: var[el].copy_to_gpu(num) for el, num in vars_remaining} - ctx.update(ctx_rest) - ctx["kv_indices"] = var["kv_indices"].gpu - kv_indices_generate_triton( - ctx["block_tables"], - ctx["kv_indices"], - ctx["kv_indptr"], - self.block_ratio, - max_seqlen_k, - ) + # All persistent metadata publishes on the owner's compute stream. + ctx = { + "kv_indptr": var["kv_indptr"].copy_to_gpu(running_bs + 1), + "block_tables": block_table_state(var["block_tables"]).publish(running_bs), + } + ctx.update({el: var[el].copy_to_gpu(num) for el, num in vars_remaining}) + ctx["kv_indices"] = var["kv_indices"].gpu + kv_indices_generate_triton( + ctx["block_tables"], + ctx["kv_indices"], + ctx["kv_indptr"], + self.block_ratio, + max_seqlen_k, + ) is_sparse_mtp = self.is_sparse and max_seqlen_q > 1 # metadata copies on main stream @@ -2692,8 +2710,6 @@ def prepare_decode( ) ctx_mla_ps_sparse = None ctx.update(ctx_mla_ps) - if not disagg: - current_stream.wait_stream(prep_stream) attn_metadata = AttentionMetaData( dropout_p=dropout_p, max_seqlen_q=max_seqlen_q, @@ -2763,6 +2779,7 @@ def _prepare_ubatch_decode( """ Splits the full-batch data into per-ubatch . """ + self._check_metadata_writable("ub0_mla_metadata", "ub1_mla_metadata") var = self.model_runner.forward_vars self._tbo_full_running_bs = bs N = self._NUM_TBO_UBATCHES @@ -2799,10 +2816,9 @@ def _prepare_ubatch_decode( ] var[f"{p}slot_mapping"].np[ub_real_tokens:ub_running_tokens] = -1 - var[f"{p}block_tables"].np[:ub_real_reqs] = var["block_tables"].np[ - req_start : req_start + ub_real_reqs - ] - var[f"{p}block_tables"].np[ub_real_reqs:running_bs] = 0 + block_table_state(var["block_tables"]).slice_to( + var[f"{p}block_tables"], req_start, ub_real_reqs, pad_to=running_bs + ) full_kv_indptr = var["kv_indptr"].np base = full_kv_indptr[req_start] @@ -2863,7 +2879,6 @@ def _prepare_ubatch_decode( (f"{p}context_lens", running_bs), (f"{p}kv_last_page_lens", running_bs), (f"{p}slot_mapping", ub_running_tokens), - (f"{p}block_tables", running_bs), (f"{p}kv_indptr", running_bs + 1), (f"{p}cu_seqlens_q", running_bs + 1), ] @@ -2872,6 +2887,7 @@ def _prepare_ubatch_decode( if self.is_sparse: vars_used.append((f"{p}sparse_kv_indptr", running_bs + 1)) + block_table_state(var[f"{p}block_tables"]).publish(running_bs) for el, num in vars_used: var[el].copy_to_gpu(num) @@ -2936,8 +2952,13 @@ def _set_ubatch_mla_buffers( max_split_per_batch=_MLA_SPLIT_BUDGET_AUTO, ) + def _capture_needs_nonempty_kv(self, max_q_len: int) -> bool: + return self.block_size > 1 or (self.dcp_world_size > 1 and max_q_len > 1) + def build_for_cudagraph_capture(self, bs: int) -> AttentionMetaData: + self._check_metadata_writable("mla_csr") var = self.model_runner.forward_vars + max_q_len = var["mtp_k"] + 1 if "mtp_k" in var else 1 self._tbo_full_running_bs = bs # Self-consistent minimal KV metadata for capture: give every sequence # exactly 1 page (kv_indptr = [0,1,...,bs]) pointing at block 0, with a @@ -2949,14 +2970,13 @@ def build_for_cudagraph_capture(self, bs: int) -> AttentionMetaData: # (only hit when num_kv_splits > 1; passes==1 takes the bf16 fast path). # Replay overwrites these buffers with real values, so this only affects # capture-time loop termination, not inference correctness. - if self.block_size > 1: + if self._capture_needs_nonempty_kv(max_q_len): kv_indptr_buf = var["kv_indptr"] kv_indptr_buf.np[: bs + 1] = np.arange(bs + 1, dtype=np.int32) kv_indptr_buf.copy_to_gpu(bs + 1) var["kv_indices"].gpu[:bs].zero_() var["kv_last_page_lens"].gpu[:bs].fill_(1) sparse_kv_indptr = var["sparse_kv_indptr"].gpu if self.is_sparse else None - max_q_len = var["mtp_k"] + 1 if "mtp_k" in var else 1 scheduled_tokens = bs * max_q_len is_sparse_mtp = self.is_sparse and max_q_len > 1 # DCP + MTP (max_q_len>1) capture: the cprr kernel masks on GLOBAL @@ -2966,17 +2986,12 @@ def build_for_cudagraph_capture(self, bs: int) -> AttentionMetaData: # with real values. cp_round_robin = self.dcp_world_size > 1 and max_q_len > 1 if cp_round_robin: - if self.block_size == 1: - var["kv_indptr"].np[: bs + 1] = np.arange(bs + 1, dtype=np.int32) - var["kv_indptr"].copy_to_gpu(bs + 1) - var["kv_indices"].gpu[:bs].zero_() - var["kv_last_page_lens"].gpu[:bs].fill_(1) # g_kv_indptr is the only thing telling the cprr kernel how long each # sequence is GLOBALLY, and nothing else initializes it -- capturing # with it left at its allocation value walks the kernel off the KV # list (illegal access). The round-robin is token-level whatever the # block size, so the one-local-token-per-rank layout set up here (or - # by the block_size > 1 branch above) is a global length of + # by the common synthetic layout above) is a global length of # dcp_world_size in both cases. var["g_kv_indptr"].np[: bs + 1] = ( np.arange(bs + 1, dtype=np.int32) * self.dcp_world_size diff --git a/atom/model_ops/attentions/backends.py b/atom/model_ops/attentions/backends.py index 587521bfef..220977426d 100644 --- a/atom/model_ops/attentions/backends.py +++ b/atom/model_ops/attentions/backends.py @@ -32,8 +32,10 @@ from atom.model_ops.attentions.token_layout.prefill import prefill_positions from atom.model_ops.attentions.token_layout.slots import slot_mapping from atom.model_ops.dcp_ops import dcp_prefill_slot_mapping -from atom.utils import CpuGpuBuffer, pack_rows +from atom.utils import CpuGpuBuffer +from atom.utils.block_tables import block_table_state from atom.utils.forward_context import AttentionMetaData, AttnState, ForwardMode +from atom.utils.h2d import h2d_producer from atom.utils.tbo.ubatch_splitting import ( UBatchSlice, attach_tbo_cpu_lens, @@ -222,7 +224,11 @@ def blank_cache_write_targets(self) -> None: A capture runs the model for real, so it writes through these; replay overwrites them with real rows. """ - for buf in self.cache_write_targets(): + targets = self.cache_write_targets() + for buf in targets: + if getattr(buf, "_publication", None) is not None: + buf._publication.acquire_write() + for buf in targets: buf.np[:] = PAD_SLOT_ID buf.copy_to_gpu() @@ -444,6 +450,8 @@ def build_kv_cache_tensor(self, module): class CommonAttentionBuilder(PoolRowsMixin, AttentionMetadataBuilder[T], Generic[T]): + capture_owns_cu_seqlens_q = False + def __init__(self, model_runner): self.model_runner = model_runner assert model_runner.block_size % self.block_size == 0 @@ -493,19 +501,24 @@ def __init__(self, model_runner): i64_kwargs = {"dtype": torch.int64, "device": self.device} i32_kwargs = {"dtype": torch.int32, "device": self.device} + prefill_i64 = dict(i64_kwargs, publication_group="prefill") + prefill_i32 = dict(i32_kwargs, publication_group="prefill") + attn_metadata = { - "slot_mapping": CpuGpuBuffer(self.max_num_batched_tokens, **i64_kwargs), - "context_lens": CpuGpuBuffer(self.max_bs, **i32_kwargs), + "slot_mapping": CpuGpuBuffer(self.max_num_batched_tokens, **prefill_i64), + "context_lens": CpuGpuBuffer(self.max_bs, **prefill_i32), "block_tables": CpuGpuBuffer( - self.max_bs, self.block_table_cols, **i32_kwargs + self.max_bs, self.block_table_cols, **prefill_i32 + ), + "cu_seqlens_q": CpuGpuBuffer( + self.max_bs + 1, **i32_kwargs, publication_group="early" ), - "cu_seqlens_q": CpuGpuBuffer(self.max_bs + 1, **i32_kwargs), - "cu_seqlens_k": CpuGpuBuffer(self.max_bs + 1, **i32_kwargs), + "cu_seqlens_k": CpuGpuBuffer(self.max_bs + 1, **prefill_i32), # Uploaded only on a prefix-cache hit, so consumers read # `AttentionMetaData.num_cached_tokens is None` as "no row has any". - "num_cached_tokens": CpuGpuBuffer(self.max_bs, **i32_kwargs), + "num_cached_tokens": CpuGpuBuffer(self.max_bs, **prefill_i32), # seq_starts for cp_mha_gather_cache: always zeros (prefix at position 0) - "seq_starts": CpuGpuBuffer(self.max_bs, **i32_kwargs), + "seq_starts": CpuGpuBuffer(self.max_bs, **prefill_i32), # token -> seq over this fwd's QUERY tokens; `-1` on the CUDAGraph # pad tail. Every backend needs it (see token_layout/batch_ids.py). "batch_id_per_q_token": CpuGpuBuffer( @@ -533,15 +546,22 @@ def _publish_indexer_fp4_decode_schedule( the target's metadata, so a write here would reach the verify step. """ - def prepare_block_tables(self, batch: ScheduledBatch): - """Marshal the batch's block tables into `forward_vars["block_tables"]`. - - Runs on every prefill step, not only the ones that upload the buffer: - `prepare_prefill` reads it back to place each token's KV slot, so a - caller that wants the table on the device only has to upload it. - """ - pack_rows(self.model_runner.forward_vars["block_tables"].np, batch.block_tables) - + def prepare_block_tables(self, batch: ScheduledBatch, running_bs=None): + """Prepare the shared CPU snapshot, reusing unchanged page mappings.""" + return block_table_state( + self.model_runner.forward_vars["block_tables"] + ).prepare(batch.block_tables, pad_to=running_bs) + + def _check_metadata_writable(self, *group_names: str) -> None: + """Preflight persistent host sources before any producer mutates them.""" + groups = getattr(self.model_runner, "h2d_groups", None) + if groups is not None: + for name in group_names: + group = groups.get(name) + if group is not None: + group.check_writable() + + @h2d_producer("mrope", runner="model_runner") def _mrope_cpu_view(self, num_tokens: int) -> np.ndarray: return ( self.model_runner.forward_vars["mrope_positions"] @@ -551,9 +571,15 @@ def _mrope_cpu_view(self, num_tokens: int) -> np.ndarray: def _copy_mrope_to_gpu(self, num_tokens: int) -> torch.Tensor: buf = self.model_runner.forward_vars["mrope_positions"] - buf.gpu.reshape(-1)[: 3 * num_tokens].copy_( - buf.cpu.reshape(-1)[: 3 * num_tokens], non_blocking=True - ) + if getattr(buf, "_publication", None) is not None: + group = self.model_runner.h2d_groups["mrope"] + group.counts[0] = 3 * num_tokens + group.publish(group.counts) + else: + # Independent, unregistered builder callers retain flat semantics. + buf.gpu.reshape(-1)[: 3 * num_tokens].copy_( + buf.cpu.reshape(-1)[: 3 * num_tokens], non_blocking=True + ) return self.model_runner._mrope_positions_view(num_tokens) def _build_mrope_prefill_positions( @@ -587,12 +613,18 @@ def _build_mrope_decode_positions( batch: ScheduledBatch, context_lens: np.ndarray, max_seqlen_q: int, + *, + running_tokens: int | None = None, ) -> torch.Tensor | None: if not getattr(self.model_runner, "use_mrope", False): return None scheduled_tokens = batch.total_tokens_num_decode - positions = self._mrope_cpu_view(scheduled_tokens) + # Every decode model view uses the padded token count as its axis + # stride, in both eager execution and graph replay. + width = scheduled_tokens if running_tokens is None else running_tokens + positions = self._mrope_cpu_view(width) + positions[:, scheduled_tokens:] = 0 offset = 0 for req_id, context_len in zip(batch.req_ids, context_lens): start = int(context_len) - max_seqlen_q @@ -605,11 +637,18 @@ def _build_mrope_decode_positions( positions[:, offset : offset + max_seqlen_q] = base[None, :] offset += max_seqlen_q - return self._copy_mrope_to_gpu(scheduled_tokens) + return self._copy_mrope_to_gpu(width)[:, :scheduled_tokens] def publish_cu_seqlens_q( self, batch: ScheduledBatch, forward_mode: ForwardMode ) -> None: + count = self.prepare_cu_seqlens_q(batch, forward_mode) + self.model_runner.forward_vars["cu_seqlens_q"].copy_to_gpu(count) + + @h2d_producer("early", runner="model_runner") + def prepare_cu_seqlens_q( + self, batch: ScheduledBatch, forward_mode: ForwardMode + ) -> int: """Publish this step's `cu_seqlens_q`. The only writer. Lives here because this class declares the buffer and defines its @@ -630,8 +669,8 @@ def publish_cu_seqlens_q( cu = self.model_runner.forward_vars["cu_seqlens_q"] cu.np[1 : scheduled_bs + 1] = np.cumsum(batch.num_scheduled_tokens) cu.np[scheduled_bs + 1 : forward_mode.running_bs + 1] = batch.total_tokens_num - # The step's only H2D for this buffer; every consumer slices `.gpu`. - cu.copy_to_gpu(forward_mode.running_bs + 1) + # Caller publishes alone or with sampling/IDs before token assembly. + return forward_mode.running_bs + 1 def decode_spans(self, batch: ScheduledBatch) -> tuple[int, np.ndarray, np.ndarray]: """A pure-decode step's `(bs, per-request lengths, exclusive prefix sum)`. @@ -729,9 +768,14 @@ def publish_batch_ids( ) -> torch.Tensor: """Build the token -> seq map into the shared buffer and upload it.""" buf = self.model_runner.forward_vars["batch_id_per_q_token"] + if getattr(buf, "_publication", None) is not None: + buf._publication.acquire_write() build_batch_ids(seqlens, pad_to=pad_to, out=buf.np) return buf.copy_to_gpu(pad_to or int(seqlens.sum())) + def _prefill_block_table_rows(self, scheduled_bs, running_bs, has_cached): + return scheduled_bs if has_cached else None + def _upload_prefill_mirrors( self, scheduled_bs: int, @@ -759,16 +803,32 @@ def _upload_prefill_mirrors( # already an int32 array rather than a Python list. var["num_cached_tokens"].np[:scheduled_bs] = cached_lens vars_used += [ - ("block_tables", scheduled_bs), ("seq_starts", scheduled_bs), ("num_cached_tokens", scheduled_bs), ] - ctx = {el: var[el].copy_to_gpu(num) for el, num in vars_used} + block_rows = self._prefill_block_table_rows( + scheduled_bs, running_bs, has_cached + ) + if block_rows is not None: + vars_used.append(("block_tables", block_rows)) + groups = self.model_runner.h2d_groups + combined = "prefill_inputs" in groups and not self.model_runner.use_mrope + group = groups["prefill_inputs" if combined else "prefill"] + counts = group.counts + for i in range(len(counts)): + counts[i] = None + for name, count in vars_used: + counts[group.indices[name]] = count + if combined: + counts[group.indices["positions"]] = scheduled_tokens + block_table_state(var["block_tables"]).publish(block_rows, group=group) + ctx = {name: var[name].gpu[:count] for name, count in vars_used} # Already on the device: `publish_cu_seqlens_q` uploads it for every # step, so this is a view. One writer AND one upload for this buffer. ctx["cu_seqlens_q"] = var["cu_seqlens_q"].gpu[: running_bs + 1] return ctx + @h2d_producer("prefill", "positions", runner="model_runner") def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): scheduled_bs = batch.total_seqs_num_prefill scheduled_tokens = batch.total_tokens_num_prefill @@ -855,6 +915,8 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): mrope_positions = self._build_mrope_prefill_positions(batch) if mrope_positions is not None: positions = mrope_positions + elif "prefill_inputs" in self.model_runner.h2d_groups: + positions = var["positions"].gpu[:scheduled_tokens] else: positions = var["positions"].copy_to_gpu(scheduled_tokens) diff --git a/atom/model_ops/attentions/deepseek_v41/backend.py b/atom/model_ops/attentions/deepseek_v41/backend.py index 355fc07a16..32ac5e42de 100644 --- a/atom/model_ops/attentions/deepseek_v41/backend.py +++ b/atom/model_ops/attentions/deepseek_v41/backend.py @@ -40,6 +40,8 @@ def get_builder_cls(): class DeepseekV41MetadataBuilder(CommonAttentionBuilder): + capture_owns_cu_seqlens_q = True + # Reuse V4's publisher and staging contract, including fixed addresses and # running_bs padding. Only pool-slot -> physical-row geometry differs. _stage = DeepseekV4AttentionMetadataBuilder._stage @@ -50,10 +52,29 @@ class DeepseekV41MetadataBuilder(CommonAttentionBuilder): # `_unique_compress_ratios_overlap` and the `v4_*_plan_{ratio}` buffers, # which the property and `__init__` below supply under V4's names. _build_compress_plans = DeepseekV4AttentionMetadataBuilder._build_compress_plans + _compress_publication_group = ( + DeepseekV4AttentionMetadataBuilder._compress_publication_group + ) # An index key is rotated at its compression group's first token, not at # its own, so the plan has to publish those positions. _publishes_key_rope = True + @property + def h2d_group_members(self): + # Plans, state slots and step rows have no GPU consumer until + # begin_step builds its indptrs. Derive members from producer groups. + return { + "v41_metadata": tuple( + name + for name, buffer in self.model_runner.forward_vars.items() + if isinstance(buffer, CpuGpuBuffer) + and ( + buffer.publication_group in ("v4_plans", "v4_state", "v41_step") + or name == "block_tables" + ) + ) + } + @staticmethod def _physical_slots(pool_slots): # V4's unified plane reverses pool slots. V4.1's EntryMajorArena uses @@ -68,6 +89,9 @@ def __init__(self, model_runner): self.max_bs, self.device, read_side=False ) ) + # The V4.1 step producer fills these together with per-ratio visibility. + for name in ("positions", "batch_id_per_q_token"): + model_runner.forward_vars[name].publication_group = "v41_step" self.config = model_runner.config.hf_config topology = build_attention_topology(self.config)[ : self.config.num_hidden_layers @@ -158,6 +182,7 @@ def _compress_plan_buffers(geometry, max_num_batched_tokens, max_bs, device): dtype=torch.int32, device=device, pin_memory=device != "cpu", + publication_group="v4_plans", ) # Sentinel, so a capture before the first real forward reads # rows the kernels skip rather than zeros -- which would name @@ -173,6 +198,7 @@ def _compress_plan_buffers(geometry, max_num_batched_tokens, max_bs, device): dtype=torch.int64, device=device, pin_memory=device != "cpu", + publication_group="v4_plans", ) # What a sentinel row works out to, so a pre-forward capture reads # the value every forward writes. @@ -194,6 +220,7 @@ def _visible_buffers(geometry, max_num_batched_tokens, device): dtype=torch.int32, device=device, pin_memory=device != "cpu", + publication_group="v41_step", ) for ratio, _ in geometry.compress_ratios } @@ -257,8 +284,9 @@ def _prepare( max_q_len=None, tentative=False, start_positions=None, + query_prefix_ready=False, ): - spans, offset, next_page = [], 0, 0 + spans, rows, offset, next_page = [], [], 0, 0 slots = batch.state_slots_committed if not batch.is_dummy_run and len(slots) != batch.total_seqs_num: raise ValueError("CSA2 requires a STATE slot for every scheduled request") @@ -273,18 +301,17 @@ def _prepare( # a live slot or PAGE, even when its fabricated block ID is 0. position = 0 count = -(-length // self.block_size) - blocks = tuple(range(next_page, next_page + count)) + blocks = np.arange(next_page, next_page + count, dtype=np.int32) next_page += count slot = len(spans) else: position = ( end - length if start_positions is None else int(start_positions[i]) ) - blocks = tuple(batch.block_tables[i]) + blocks = batch.block_tables[i] slot = slots[i] - spans.append( - RequestSpan(request_id, position, offset, length, slot, blocks) - ) + spans.append(RequestSpan(request_id, position, offset, length, slot)) + rows.append(blocks) offset += length if offset != batch.total_tokens_num or running_tokens < offset: raise ValueError("CSA2 batch token spans disagree with the runner") @@ -301,12 +328,20 @@ def _prepare( ) if cache is None: raise RuntimeError("CSA2 cache must be allocated before serving") + groups = getattr(self.model_runner, "h2d_groups", None) + combined = None if groups is None else groups.get("v41_metadata") + if combined is not None and combined.transport != "packed": + combined = None + if combined is not None: + for i in range(len(combined.counts)): + combined.counts[i] = None # Zero-token scheduler rows are excluded from spans. Publish in this # same request order, including the private dummy slots used at startup. state_slot_out = self._populate_state_slot_mappings( SimpleNamespace(state_slots_committed=[span.slot for span in spans]), len(spans), running_bs, + publication_group=combined, ) verifying = tentative and not batch.is_dummy_run and bool(spans) # One plan per ratio for the whole batch, into the fixed-address @@ -322,9 +357,11 @@ def _prepare( running_bs=None if max_q_len is None else running_bs, max_q_len=max_q_len, extra_write=self.geometry.speculative_tokens if verifying else 0, + defer_to=combined, ) step = cache.begin_step( spans, + block_tables=rows, tentative=verifying, buffers=self.model_runner.forward_vars, running_bs=running_bs, @@ -332,6 +369,17 @@ def _prepare( max_q_len=max_q_len, state_slot_out=state_slot_out, plans=plans, + publication_group=( + combined + if combined is not None + else None if groups is None else groups["v41_step"] + ), + query_prefix_ready=query_prefix_ready and len(spans) == len(batch.req_ids), + query_prefix_republish_reason=( + "compact zero-token scheduler rows for CSA2 after input assembly" + if query_prefix_ready and len(spans) != len(batch.req_ids) + else None + ), ) positions = self.model_runner.forward_vars["positions"] cu = self.model_runner.forward_vars["cu_seqlens_q"].gpu[: running_bs + 1] @@ -365,7 +413,9 @@ def _prepare( return metadata, positions.gpu[:running_tokens] def prepare_prefill(self, batch, running_bs): - return self._prepare(batch, running_bs, batch.total_tokens_num) + return self._prepare( + batch, running_bs, batch.total_tokens_num, query_prefix_ready=True + ) def prepare_decode(self, batch, running_bs, running_tokens, max_seqlen_q): starts = None @@ -383,6 +433,7 @@ def prepare_decode(self, batch, running_bs, running_tokens, max_seqlen_q): max_q_len=max_seqlen_q, tentative=bool(self.geometry.speculative_tokens), start_positions=starts, + query_prefix_ready=True, ) def _engram_batch(self, step, cache, metadata, tokens): @@ -446,6 +497,14 @@ def prepare_model_inputs(self, input_ids, metadata): # stages addresses only -- the slots are somebody else's, and # `prepare_state`'s position-0 reset would zero their state. batch = self._engram_batch(step, cache, metadata, tokens) + # Tentative candidates have separate storage: freeze the hash inputs + # and stage every accepted prefix with one kernel. Committed cursors + # may alias history, so that path still advances AFTER the snapshot. + stage_cursor = ( + batch is not None + and step.tentative + and self.engram.host.overlap is not None + ) histories = ( np.full((step.scheduled_bs, self.geometry.history_size), -1, np.int64) if metadata.dummy @@ -460,6 +519,12 @@ def prepare_model_inputs(self, input_ids, metadata): token_mask=metadata.token_mask, padded_rows=step.width, batch=batch, + cursor_positions=step.positions if stage_cursor else None, + cursor_out=( + cache.tentative_staging.gpu[: step.scheduled_bs] + if stage_cursor + else None + ), ) embeddings, histories = prepared.embeddings, prepared.histories if batch is None and cache.pending is not None: @@ -486,7 +551,9 @@ def prepare_model_inputs(self, input_ids, metadata): # overwrites -- the hash kernel, the staging above -- has already run. # A tentative step's cursor is the sampler's to write, so the device # path stages its candidates here and commits them there. - if batch is not None: + if stage_cursor: + cache.pending.staged_on_device = True + elif batch is not None: self._write_engram_cursor(step, cache, batch) elif not metadata.dummy and not step.tentative: cache.advance_cursor(step, histories) diff --git a/atom/model_ops/attentions/deepseek_v41/cache.py b/atom/model_ops/attentions/deepseek_v41/cache.py index 4e5819a09b..8d96ccf63b 100644 --- a/atom/model_ops/attentions/deepseek_v41/cache.py +++ b/atom/model_ops/attentions/deepseek_v41/cache.py @@ -211,6 +211,7 @@ def begin_step( self, requests, *, + block_tables, tentative=False, buffers=None, running_bs=None, @@ -218,6 +219,9 @@ def begin_step( max_q_len=None, state_slot_out=None, plans=None, + publication_group=None, + query_prefix_ready=False, + query_prefix_republish_reason=None, ): self.require_committed() requests = tuple(requests) @@ -234,7 +238,9 @@ def begin_step( ) offset = 0 seen = set() - for span in requests: + if len(block_tables) != len(requests): + raise ValueError("One PAGE table is required per request") + for span, row in zip(requests, block_tables): if span.length <= 0 or span.position < 0 or span.offset != offset: raise ValueError( "Request spans must be nonempty and partition the token batch" @@ -242,15 +248,15 @@ def begin_step( if not 0 <= span.slot < self.num_slots or span.slot in seen: raise ValueError("Each request needs its own valid STATE slot") needed = -(-span.end // self.geometry.block_size) - if len(span.block_ids) < needed or any( - block < 0 or block >= self.num_pages for block in span.block_ids - ): + if len(row) < needed: raise ValueError("Request PAGE table is incomplete or out of range") seen.add(span.slot) offset += span.length step = prepare_batch_step( requests, self.pool.device, + block_tables=block_tables, + page_limit=self.num_pages, tentative=tentative, buffers=buffers, running_bs=running_bs, @@ -258,6 +264,9 @@ def begin_step( max_q_len=max_q_len, state_slot_out=state_slot_out, ratios=tuple(ratio for ratio, _ in self.geometry.compress_ratios), + publication_group=publication_group, + query_prefix_ready=query_prefix_ready, + query_prefix_republish_reason=query_prefix_republish_reason, ) step.plans = ( self._private_plans(requests, tentative) if plans is None else plans diff --git a/atom/model_ops/attentions/deepseek_v41/metadata.py b/atom/model_ops/attentions/deepseek_v41/metadata.py index 1e753e2288..b78c525623 100644 --- a/atom/model_ops/attentions/deepseek_v41/metadata.py +++ b/atom/model_ops/attentions/deepseek_v41/metadata.py @@ -8,7 +8,8 @@ from atom.model_ops.attentions.token_layout.batch_ids import build_batch_ids from atom.model_ops.attentions.token_layout.prefill import prefill_positions -from atom.utils import CpuGpuBuffer, pack_rows +from atom.utils import CpuGpuBuffer +from atom.utils.block_tables import block_table_state @dataclass(frozen=True) @@ -18,7 +19,6 @@ class RequestSpan: offset: int length: int slot: int - block_ids: tuple[int, ...] @property def end(self): @@ -113,31 +113,12 @@ def visible_buffer_name(ratio): return f"v41_index_visible_{ratio}" -def _publish_block_tables(tables, requests, running_bs): - """Reuse the fixed device page table until its rows or padding change. - - Every V4.1 publication, including dummy/capture batches, goes through here. - DSpark borrows padding entries only temporarily and restores them before - returning. Cache the mapping, not positions: advancing within an allocated - page does not change an address in this table. - """ - rows = tuple(tuple(span.block_ids) for span in requests) - key = (running_bs, rows) - previous = getattr(tables, "_v41_published_rows", None) - if previous is not None and previous[0] is tables.gpu and previous[1] == key: - return tables.gpu[:running_bs] - if rows: - pack_rows(tables.np, [np.asarray(row, dtype=np.int32) for row in rows]) - tables.np[len(rows) : running_bs] = 0 - published = tables.copy_to_gpu(running_bs) - tables._v41_published_rows = (tables.gpu, key) - return published - - def prepare_batch_step( requests, device, *, + block_tables, + page_limit=None, tentative=False, buffers=None, running_bs=None, @@ -145,6 +126,9 @@ def prepare_batch_step( max_q_len=None, state_slot_out=None, ratios=(), + publication_group=None, + query_prefix_ready=False, + query_prefix_republish_reason=None, ): """Stage request metadata using the same persistent buffers/layout as V4. @@ -164,7 +148,7 @@ def prepare_batch_step( if max_q_len is not None and lengths.size and max_q_len < int(lengths.max()): raise ValueError("A request is longer than the query width this forward runs") if buffers is None: - width = max((len(span.block_ids) for span in requests), default=0) + width = max((len(row) for row in block_tables), default=0) shapes = { "positions": (running_tokens,), "cu_seqlens_q": (running_bs + 1,), @@ -191,10 +175,22 @@ def prepare_batch_step( for name, count in required.items(): if count > buffers[name].np.shape[0]: raise ValueError(f"{name} metadata buffer cannot hold {count} rows") + if publication_group is not None: + publication_group.check_writable() + tables = block_table_state(buffers["block_tables"]).prepare( + block_tables, pad_to=running_bs, page_limit=page_limit + ) cu = buffers["cu_seqlens_q"] - cu.np[0] = 0 - np.cumsum(lengths, out=cu.np[1 : scheduled_bs + 1]) - cu.np[scheduled_bs + 1 : running_bs + 1] = scheduled_tokens + if not query_prefix_ready: + if cu._publication is not None: + # Reject unannounced rewrites before touching the pinned prefix. + # Explicit compaction may reacquire it after token assembly. + cu._publication.acquire_write( + republish_reason=query_prefix_republish_reason + ) + cu.np[0] = 0 + np.cumsum(lengths, out=cu.np[1 : scheduled_bs + 1]) + cu.np[scheduled_bs + 1 : running_bs + 1] = scheduled_tokens positions = buffers["positions"] starts = np.asarray([span.position for span in requests], dtype=positions.np.dtype) prefill_positions( @@ -225,14 +221,36 @@ def prepare_batch_step( dtype=torch.int32, device=device, ) - published = { - name: buffers[name].copy_to_gpu(count) - for name, count in required.items() - if name != "block_tables" - } - published["block_tables"] = _publish_block_tables( - buffers["block_tables"], requests, running_bs + if publication_group is not None: + grouped_tables = "block_tables" in publication_group.indices + for i, member in enumerate(publication_group.members): + if member.name == "block_tables": + publication_group.counts[i] = running_bs + elif member.name in required: + publication_group.counts[i] = required[member.name] + # Plans and state slots were staged by the builder. Preserve + # their counts in this combined publication, before indptrs run. + tables.publish(running_bs if grouped_tables else None, group=publication_group) + published = { + member.name: member.destination[: required[member.name]] + for member in publication_group.members + if member.name in required + } + else: + published = { + name: buffers[name].copy_to_gpu(count) + for name, count in required.items() + if name not in ("block_tables", "cu_seqlens_q") + } + published["cu_seqlens_q"] = ( + cu.gpu[: running_bs + 1] + if query_prefix_ready + else cu.copy_to_gpu( + running_bs + 1, republish_reason=query_prefix_republish_reason + ) ) + if "block_tables" not in published: + published["block_tables"] = tables.publish(running_bs) return BatchStep( requests, published["positions"], diff --git a/atom/model_ops/attentions/deepseek_v4_attn.py b/atom/model_ops/attentions/deepseek_v4_attn.py index 0531d0a94d..fb5214d4f2 100644 --- a/atom/model_ops/attentions/deepseek_v4_attn.py +++ b/atom/model_ops/attentions/deepseek_v4_attn.py @@ -41,6 +41,7 @@ import math import os from collections.abc import Sequence +from contextlib import contextmanager from dataclasses import dataclass from typing import Any, cast @@ -122,12 +123,14 @@ write_v4_paged_prefill_indices, ) from atom.utils import CpuGpuBuffer, envs, upload_numpy +from atom.utils.block_tables import block_table_state from atom.utils.forward_context import ( AttentionMetaData, AttnState, Context, get_forward_context, ) +from atom.utils.h2d import h2d_producer _FP4_OPUS_DECODE_Q1_VARIANT = "qlen1_kv64" _FP4_OPUS_DECODE_Q4_VARIANT = "qlen4_kv64" @@ -433,6 +436,7 @@ class DeepseekV4AttentionMetadataBuilder(CommonAttentionBuilder): `config.kv_cache_block_size` (config.py forces the same value for V4). """ + capture_owns_cu_seqlens_q = True block_size = 256 # Number of micro-batches for Two-Batch Overlap (TBO). @@ -583,6 +587,7 @@ def __init__(self, model_runner): self.block_table_cols, dtype=torch.int32, device=self.device, + publication_group="prefill", ) # What one CSA layer's indexer block holds, and where each region sits # in it. Sizing, allocation and the scale view `build_kv_cache_tensor` @@ -2140,6 +2145,7 @@ def _build_fp4_opus_prefill_plans( cu_seqlens_q_cpu: np.ndarray, reuse_cu_seqlens_q: bool, plan_total_tokens: int | None, + buf_prefix_ubatch: str = "", ) -> None: from aiter.ops.opus.pa_mqa_logits_mxfp4 import pa_mqa_logits_mxfp4_plan @@ -2156,21 +2162,47 @@ def _build_fp4_opus_prefill_plans( ) row_to_batch = attn_metadata.batch_id_per_q_token[:total_tokens] starts = torch.zeros_like(visible_end_gpu) - plans = [] + chunks = [] + rows = 0 for chunk_start in range(0, plan_total, chunk_tokens): chunk_end = min(chunk_start + chunk_tokens, plan_total) cu_chunk_np = _chunk_cu_seqlens(cu_cpu, chunk_start, chunk_end) - if ( + reuse = ( reuse_cu_seqlens_q and chunk_start == 0 and chunk_end == total_tokens and cu_chunk_np.size == cu_cpu.size - ): + ) + chunks.append((chunk_start, chunk_end, cu_chunk_np, reuse)) + if not reuse: + rows += cu_chunk_np.size + + chunk_gpu = None + if rows: + buf = self.model_runner.forward_vars[ + f"{buf_prefix_ubatch}v4_indexer_chunk_cu" + ] + if rows > buf.cpu.numel(): + raise ValueError("Opus chunk prefixes exceed initialized capacity") + if buf._publication is not None: + buf._publication.acquire_write() + at = 0 + for _, _, cu_chunk_np, reuse in chunks: + if not reuse: + end = at + cu_chunk_np.size + buf.np[at:end] = cu_chunk_np + at = end + chunk_gpu = buf.copy_to_gpu(rows) + + plans = [] + at = 0 + for chunk_start, chunk_end, cu_chunk_np, reuse in chunks: + if reuse: cu_chunk = attn_metadata.cu_seqlens_q[: cu_cpu.size] else: - cu_chunk = torch.from_numpy(cu_chunk_np).to( - visible_end_gpu.device, non_blocking=True - ) + end = at + cu_chunk_np.size + cu_chunk = chunk_gpu[at:end] + at = end plan = pa_mqa_logits_mxfp4_plan( cu_chunk, visible_end_gpu[chunk_start:chunk_end], @@ -2384,6 +2416,7 @@ def _build_v4_indexer_meta( cu_seqlens_q_cpu=cu_seqlens_q_cpu, reuse_cu_seqlens_q=reuse_cu_seqlens_q, plan_total_tokens=opus_plan_total_tokens, + buf_prefix_ubatch=buf_prefix_ubatch, ) return meta # Precompute the FP4 prefill persistent-grid schedule here (instead @@ -2577,12 +2610,12 @@ def prepare_mtp_decode( attn_metadata.batch_id_per_q_token = batch_id_per_q_token # fp8 asm decode per-token index tensors. MTP draft step is 1-token-per- - # seq → the asm kernel sees N = bs. Stage the constant per-token tensors - # to that length via the same builder-staged path as the verify fwd. + # seq → the asm kernel sees N = bs. Its prefix is an immutable arange; + # verify's padded prefix may repeat its last token and stays separate. if self._kv_fp8: - attn_metadata.qo_indptr = self._stage( - "v4_qo_indptr", self._v4_qo_indptr_np[: running_bs + 1] - ) + # Draft rows are uniformly one token each. This immutable device + # prefix avoids rewriting verify's pinned source mid-forward. + attn_metadata.qo_indptr = var["v4_draft_qo_indptr"][: running_bs + 1] attn_metadata.empty_kv_indptr = self.model_runner.forward_vars[ "v4_empty_kv_indptr" ][: running_bs + 1] @@ -2596,6 +2629,7 @@ def prepare_mtp_decode( # - v4 indexer meta (Indexer — only present when ratio == 4) return {} + @h2d_producer("prefill", "positions", "v4_state", runner="model_runner") def prepare_decode( self, batch: ScheduledBatch, @@ -2606,9 +2640,8 @@ def prepare_decode( """V4-style decode prep: populates positions, cu_seqlens_q, block_tables, and state_slot_out. - Uses stream overlap (like AiterMLAMetadataBuilder) to hide H2D - latency behind CPU numpy work: basic H2D copies fire on - ``prep_stream`` while ``_build_compress_plans`` runs on the CPU. + Publishes metadata on the current compute stream. Host sources remain + borrowed until the runner's preparation completion event. """ var = self.model_runner.forward_vars scheduled_bs, lens, cu = self.decode_spans(batch) @@ -2670,7 +2703,7 @@ def prepare_decode( var["context_lens"].np[:scheduled_bs] = context_lens_np - # Inline block_tables CPU fill (H2D deferred to prep_stream). + # Fill block tables before their current-stream publication. self.prepare_block_tables(batch) pool_np = np.asarray(batch.state_slots_committed[:scheduled_bs], dtype=np.int32) @@ -2692,22 +2725,16 @@ def prepare_decode( ss_buf.np[scheduled_bs:running_bs] = 0 si_buf.np[scheduled_bs:running_bs] = 0 - # ---- fire H2D on prep_stream ---- - # NB: this runs inside attn_metadata_builder.build(), BEFORE - # set_forward_context() — can't read main_stream from the context yet. - prep_stream = self.prep_stream - current_stream = torch.cuda.current_stream() - prep_stream.wait_stream(current_stream) - with torch.cuda.stream(prep_stream): - positions = var["positions"].copy_to_gpu(running_tokens) - # Uploaded once by `publish_cu_seqlens_q`; this is a view. - cu_seqlens_q_gpu = var["cu_seqlens_q"].gpu[: running_bs + 1] - context_lens_gpu = var["context_lens"].copy_to_gpu(scheduled_bs) - block_tables_gpu = var["block_tables"].copy_to_gpu(scheduled_bs) - state_slot_gpu = ss_buf.copy_to_gpu(running_bs) - state_slot_in_gpu = si_buf.copy_to_gpu(running_bs) - - # ---- CPU numpy work, overlapped with prep_stream H2D ---- + # Publish before the first GPU metadata consumer. + positions = var["positions"].copy_to_gpu(running_tokens) + # Uploaded once by `publish_cu_seqlens_q`; this is a view. + cu_seqlens_q_gpu = var["cu_seqlens_q"].gpu[: running_bs + 1] + context_lens_gpu = var["context_lens"].copy_to_gpu(scheduled_bs) + block_tables_gpu = block_table_state(var["block_tables"]).publish(scheduled_bs) + state_slot_gpu = ss_buf.copy_to_gpu(running_bs) + state_slot_in_gpu = si_buf.copy_to_gpu(running_bs) + + # Build the remaining metadata on the CPU. # The plan wants the seq length THROUGH this step's last forwarded # token, not the reservation: the unreduced `ctx` tail-anchors it and # lands compressed rows `full_q - len_i` off from `visible_csa(pos)`. @@ -2722,8 +2749,7 @@ def prepare_decode( extra_write=self.max_spec_steps, ) - # ---- sync, build attn_metadata, per-fwd meta ---- - current_stream.wait_stream(prep_stream) + # Build attention metadata after its uploads have been submitted. attn_metadata = AttentionMetaData_DSV4( cu_seqlens_q=cu_seqlens_q_gpu, @@ -2802,6 +2828,8 @@ def _prepare_ubatch_decode( boundaries for DSpark ragged verify; rectangular decode contains the uniform ``max_seqlen_q`` value. """ + for ub_idx in range(self._NUM_TBO_UBATCHES): + self._check_ubatch_sources(f"ub{ub_idx}_") var = self.model_runner.forward_vars N = self._NUM_TBO_UBATCHES enforce_eager = self.model_runner.enforce_eager @@ -2869,10 +2897,9 @@ def _prepare_ubatch_decode( var[f"{p}v4_meta_state_slot_in"].np[:ub_real_reqs] = ub_state_in_np var[f"{p}v4_meta_state_slot_in"].np[ub_real_reqs:ub_running_bs] = 0 - var[f"{p}block_tables"].np[:ub_real_reqs] = var["block_tables"].np[ - req_start : req_start + ub_real_reqs - ] - var[f"{p}block_tables"].np[ub_real_reqs:ub_running_bs] = 0 + block_table_state(var["block_tables"]).slice_to( + var[f"{p}block_tables"], req_start, ub_real_reqs, pad_to=ub_running_bs + ) # positions: copy the ubatch's token slice (values match the global # positions slice the UBatchWrapper Context will expose). @@ -2894,7 +2921,9 @@ def _prepare_ubatch_decode( positions_gpu = var[f"{p}positions"].copy_to_gpu(ub_running_tokens) cu_seqlens_q_gpu = var[f"{p}cu_seqlens_q"].copy_to_gpu(ub_running_bs + 1) context_lens_gpu = var[f"{p}context_lens"].copy_to_gpu(ub_running_bs) - block_tables_gpu = var[f"{p}block_tables"].copy_to_gpu(ub_running_bs) + block_tables_gpu = block_table_state(var[f"{p}block_tables"]).publish( + ub_running_bs + ) state_slot_gpu = var[f"{p}v4_meta_state_slot_out"].copy_to_gpu( ub_running_bs ) @@ -2989,9 +3018,9 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): if attn_metadata.block_tables is None: # Marshalled by the parent every step; only its upload is gated on # a prefix-cache hit, and V4 needs the table on every one. - attn_metadata.block_tables = self.model_runner.forward_vars[ - "block_tables" - ].copy_to_gpu(scheduled_bs) + attn_metadata.block_tables = block_table_state( + self.model_runner.forward_vars["block_tables"] + ).publish(scheduled_bs) state_slot_gpu, state_slot_np = self._populate_state_slot_mappings( batch, scheduled_bs, running_bs, return_cpu=True ) @@ -3045,15 +3074,22 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): scheduled_bs, scheduled_tokens, ) - self._attach_v4_indexer_meta( - attn_metadata, - scheduled_bs, - scheduled_tokens, - positions_gpu=positions, - cu_seqlens_q_cpu=( - cu_seqlens_q_np if self.indexer_layout == FP4_GFX1250_NATURAL else None - ), - ) + if pcp_is_enabled() and not batch.is_dummy_run: + # Non-None requests a single indexer build after query reindex. + # No full-query indexer has a consumer on this branch. + attn_metadata.indexer_meta = {} + else: + self._attach_v4_indexer_meta( + attn_metadata, + scheduled_bs, + scheduled_tokens, + positions_gpu=positions, + cu_seqlens_q_cpu=( + cu_seqlens_q_np + if self.indexer_layout == FP4_GFX1250_NATURAL + else None + ), + ) # Two-source paged_prefill index buffers (extend + per-ratio prefix). # Eager-only — direct H2D, no forward_vars staging required. Sets # attn_metadata.{kv_indices,kv_indptr}_{extend,prefix_swa,prefix_csa,prefix_hca} @@ -3116,6 +3152,7 @@ def _apply_pcp_reindex( scheduled_bs: int, total_tokens: int, cu_seqlens_q_cpu: np.ndarray | None = None, + buf_prefix_ubatch: str = "", ) -> torch.Tensor: """Reduce per-query prefill metadata to this PCP rank's round-robin shard. @@ -3142,7 +3179,10 @@ def _apply_pcp_reindex( padded_total = pcp_pad_len(total_tokens, pcp_size) n_pad = padded_total - total_tokens owned_q_cpu = pcp_round_robin_query_indices(padded_total, pcp_size) - owned_q = owned_q_cpu.to(device) + # Arithmetic metadata: generate directly on the GPU, without a + # temporary host tensor upload. Keep the CPU indices for local lengths. + first = int(owned_q_cpu[0]) if owned_q_cpu.numel() else 0 + owned_q = torch.arange(first, padded_total, pcp_size, device=device) # --- ragged per-query buffers: pad indptr to padded_total, then 1/W --- for ind_attr, idx_attr in ( @@ -3210,6 +3250,7 @@ def _apply_pcp_reindex( cu_seqlens_q_cpu=local_cu_cpu, reuse_cu_seqlens_q=False, opus_plan_total_tokens=local_tokens - dummy_rows, + buf_prefix_ubatch=buf_prefix_ubatch, ) return positions_local @@ -3217,35 +3258,62 @@ def _get_ubatch_compress_plan_buffers( self, ubatch_idx: int ) -> dict[int, dict[str, "CpuGpuBuffer"]]: - if not hasattr(self, "_ubatch_compress_plan_buffers"): - self._ubatch_compress_plan_buffers: dict[ - int, dict[int, dict[str, CpuGpuBuffer]] - ] = {} - cached = self._ubatch_compress_plan_buffers.get(ubatch_idx) - if cached is not None: - return cached - + # Both decode and prefill use the producer-declared, ring-cloned + # ubatch buffers. No first-forward pinned allocation or private alias + # outside the runner's source reuse gate. var = self.model_runner.forward_vars - pool: dict[int, dict[str, CpuGpuBuffer]] = {} - for ratio, _ in self._unique_compress_ratios_overlap: - tmpl_c = var[f"v4_compress_plan_{ratio}"] - tmpl_w = var[f"v4_write_plan_{ratio}"] - buf_c = CpuGpuBuffer( - *tmpl_c.cpu.shape, dtype=tmpl_c.cpu.dtype, device=tmpl_c.gpu.device - ) - buf_w = CpuGpuBuffer( - *tmpl_w.cpu.shape, dtype=tmpl_w.cpu.dtype, device=tmpl_w.gpu.device - ) - # Sentinel-fill so any unused tail rows behave like the main pool. - buf_c.cpu.fill_(-1) - buf_c.copy_to_gpu() - buf_w.cpu.fill_(-1) - buf_w.copy_to_gpu() - pool[ratio] = {"compress": buf_c, "write": buf_w} - self._ubatch_compress_plan_buffers[ubatch_idx] = pool - return pool + prefix = f"ub{ubatch_idx}_" + return { + ratio: { + "compress": var[f"{prefix}v4_compress_plan_{ratio}"], + "write": var[f"{prefix}v4_write_plan_{ratio}"], + } + for ratio, _ in self._unique_compress_ratios_overlap + } + + @contextmanager + def _ubatch_prefill_sources(self, ubatch_idx): + owner = getattr(self.model_runner, "h2d_owner", None) + if owner is None: + yield + return + resumed = owner._state == "sealed" + if resumed: + # Reject duplicate/binding errors before reopening a sealed phase. + self._check_ubatch_sources(f"ub{ubatch_idx}_", sealed=True) + owner.resume() + else: + self._check_ubatch_sources(f"ub{ubatch_idx}_") + try: + yield + except BaseException: + owner.fail() + raise + else: + if resumed: + owner.finish() + + def _check_ubatch_sources(self, prefix, *, sealed=False): + groups = getattr(self.model_runner, "h2d_groups", None) + if groups is not None: + for suffix in ("v4_inputs", "v4_state", "v4_indexer", "v4_plans"): + group = groups.get(f"{prefix}{suffix}") + if group is not None: + if sealed: + for member in group.members: + member._validate(0, None) + else: + group.check_writable() def build_ubatch_prefill_metadata( + self, attn_metadata, ub_slice, running_bs, ubatch_idx=0 + ) -> AttentionMetaData_DSV4: + with self._ubatch_prefill_sources(ubatch_idx): + return self._build_ubatch_prefill_metadata( + attn_metadata, ub_slice, running_bs, ubatch_idx + ) + + def _build_ubatch_prefill_metadata( self, attn_metadata: AttentionMetaData, ub_slice, @@ -3317,6 +3385,7 @@ def build_ubatch_prefill_metadata( np.ascontiguousarray(context_lens_np, dtype=np.int32), self._unique_compress_ratios_overlap, plan_buffers=ub_plan_buffers, + publication_group=self._compress_publication_group(f"ub{ubatch_idx}_"), extra_write=0, # TBO prefill is eager-only; nothing rejects it. ) else: @@ -3328,6 +3397,12 @@ def build_ubatch_prefill_metadata( # context_lens needs staging into the prefixed set here. p = f"ub{ubatch_idx}_" var[f"{p}context_lens"].np[:ub_num_reqs] = context_lens_np + # A request may straddle the token split. Publish its clamped prefix + # before indexer/paged-prefill consumers; reuse it for attention too. + ub_cu_gpu = self._stage(f"{p}cu_seqlens_q", ub_cu) + ub_attn.cu_seqlens_q = ub_cu_gpu + if ub_attn.cu_seqlens_k is not None: + ub_attn.cu_seqlens_k = ub_cu_gpu self._attach_v4_per_fwd_meta( ub_attn, @@ -3347,7 +3422,7 @@ def build_ubatch_prefill_metadata( cu_seqlens_q_cpu=( ub_cu if self.indexer_layout == FP4_GFX1250_NATURAL else None ), - reuse_cu_seqlens_q=False, + reuse_cu_seqlens_q=True, buf_prefix_ubatch=p, ) @@ -3355,9 +3430,7 @@ def build_ubatch_prefill_metadata( ub_start_pos_per_seq_np = positions_np[ub_cu[:ub_num_reqs]] ub_positions_gpu = var["positions"].gpu[ts.start : ts.stop] ub_block_tables_gpu = var["block_tables"].gpu[rs.start : rs.stop] - ub_cu_q_per_seq_gpu = torch.from_numpy( - np.ascontiguousarray(ub_cu[:ub_num_reqs], dtype=np.int32) - ).to(self.device, non_blocking=True) + ub_cu_q_per_seq_gpu = ub_cu_gpu[:ub_num_reqs] self._build_paged_prefill_meta( ub_attn, positions_np, @@ -3372,35 +3445,11 @@ def build_ubatch_prefill_metadata( block_tables_gpu=ub_block_tables_gpu, ) - # `split_attn_metadata` computed ub_attn.cu_seqlens_q/k from RAW request - # boundaries (orig_cu[rs] - base), which is WRONG for a straddling - # request under token-midpoint splits: it counts the request's FULL - # length instead of only the portion owned by this ubatch, so - # cu_seqlens_q[-1] > ub_num_tokens and any kernel indexing by it goes - # out of bounds (SIGABRT / GPU memory fault). Overwrite with the - # token-window-clamped `ub_cu` already computed above. For non- - # straddling splits these are identical, so this is a no-op there. - ub_cu_gpu = torch.from_numpy( - np.ascontiguousarray(ub_cu[: ub_num_reqs + 1], dtype=np.int32) - ).to(self.device, non_blocking=True) - ub_attn.cu_seqlens_q = ub_cu_gpu if extend_lens_np.size > 0: ub_attn.max_seqlen_q = int(extend_lens_np.max()) - # cu_seqlens_k consistent with the clamped q lens (V4 prefill prefix KV - # is read via per-ratio kv_indices_prefix_* buffers, not cu_seqlens_k). - if ub_attn.cu_seqlens_k is not None: - ub_attn.cu_seqlens_k = ub_cu_gpu - - # Clone all GPU tensors that are views into shared CpuGpuBuffers. - # Without this, building the next ubatch overwrites this ubatch's - # data via the same underlying buffer. - if ub_attn.batch_id_per_q_token is not None: - ub_attn.batch_id_per_q_token = ub_attn.batch_id_per_q_token.clone() - if ub_attn.indexer_meta is not None: - im = ub_attn.indexer_meta - if im.get("cu_committed_gpu") is not None: - im["cu_committed_gpu"] = im["cu_committed_gpu"].clone() + # Each ubatch owns these staging buffers through its runner slot. + # Consumers can use their stable views without an extra device copy. return ub_attn def _build_ubatch_prefill_metadata_balanced( @@ -3483,19 +3532,17 @@ def _build_ubatch_prefill_metadata_balanced( setattr(ub, idx_attr, nx) # batch_id_per_q_token: slice + rebase global req id → group-local (keep -1). if src.batch_id_per_q_token is not None: - bid = src.batch_id_per_q_token[gts:gte].clone() + bid = src.batch_id_per_q_token[gts:gte] ub.batch_id_per_q_token = torch.where(bid >= 0, bid - rs0, bid) if src.skip_prefix_len_csa is not None: ub.skip_prefix_len_csa = src.skip_prefix_len_csa[gts:gte].contiguous() ub.envelope_rows = src.envelope_rows - # ---- compress_plans: group's GLOBAL per-request (compressor all-gathers - # the group to full order). Built from global cu / context_lens slices. ---- + # The query lengths are already known on the host. Reuse them for + # plans, Opus prefixes and max_seqlen_q without a GPU reduction/readback. + gcu = var["cu_seqlens_q"].np + ext = (gcu[rs0 + 1 : rs1 + 1] - gcu[rs0:rs1]).astype(np.int32) if self._unique_compress_ratios_overlap: - gcu = var[ - "cu_seqlens_q" - ].np # GLOBAL (not overwritten for request-boundary split) - ext = (gcu[rs0 + 1 : rs1 + 1] - gcu[rs0:rs1]).astype(np.int32) ctx = np.asarray(var["context_lens"].np[rs0:rs1], dtype=np.int32) plan_bufs = self._get_ubatch_compress_plan_buffers(ubatch_idx) ub.compress_plans = make_compress_plans( @@ -3503,6 +3550,7 @@ def _build_ubatch_prefill_metadata_balanced( np.ascontiguousarray(ctx, dtype=np.int32), self._unique_compress_ratios_overlap, plan_buffers=plan_bufs, + publication_group=self._compress_publication_group(f"ub{ubatch_idx}_"), decode_capacity_per_ratio=None, extra_write=self.max_spec_steps, ) @@ -3524,21 +3572,12 @@ def _build_ubatch_prefill_metadata_balanced( group_bs, group_total, cu_seqlens_q_cpu=group_cu, + buf_prefix_ubatch=f"ub{ubatch_idx}_", ) - # max_seqlen_q from the group's per-request extend lengths. - if ub.cu_seqlens_q is not None and group_bs > 0: - per_req_q = ub.cu_seqlens_q[1 : group_bs + 1] - ub.cu_seqlens_q[:group_bs] - if per_req_q.numel() > 0: - ub.max_seqlen_q = int(per_req_q.max().item()) - - # Clone GPU tensors that are slices/views into shared CpuGpuBuffers, so a - # later ubatch (or fwd) reusing the same buffer can't overwrite this - # ubatch's data (mirrors the token-split path's clones). - if ub.indexer_meta is not None: - im = ub.indexer_meta - if im.get("cu_committed_gpu") is not None: - im["cu_committed_gpu"] = im["cu_committed_gpu"].clone() + if group_bs > 0: + ub.max_seqlen_q = int(ext.max()) + # cu_committed_gpu already belongs to this ubatch's runner slot. return ub def _attach_v4_per_fwd_meta( @@ -3884,11 +3923,16 @@ def _attach_v4_paged_decode_meta( # exactly this, so a temporary to hand `_stage` would be a second pass # over the token axis. `_stage`'s capacity check is owed here instead. if self._kv_fp8: - qo_buf = self.model_runner.forward_vars["v4_qo_indptr"] + # The global verify batch and each TBO ubatch have different live + # token counts. Keep their padded prefixes in independent storage. + qo_name = f"{buf_prefix_ubatch}v4_qo_indptr" + qo_buf = var[qo_name] assert T_pad + 1 <= qo_buf.np.shape[0], ( - f"V4 buffer 'v4_qo_indptr' too small: need {T_pad + 1}, have " + f"V4 buffer {qo_name!r} too small: need {T_pad + 1}, have " f"{qo_buf.np.shape[0]}. Increase T_dec in _alloc_v4_metadata_buffers." ) + if qo_buf._publication is not None: + qo_buf._publication.acquire_write() qo_buf.np[: T + 1] = self._v4_qo_indptr_np[: T + 1] qo_buf.np[T + 1 : T_pad + 1] = T attn_metadata.qo_indptr = qo_buf.copy_to_gpu(T_pad + 1) @@ -4014,9 +4058,9 @@ def _build_paged_prefill_meta( hca_total = int(hca_indptr_np[T]) # ----- H2D: 4 indptrs + 2 per-seq scalars ----- - # Sources are per-call temp np arrays, so not a cross-ubatch race source - # (the shared-pinned-buffer race is handled by the stream sync before - # build_ubatch_prefill_metadata's finally). Via `upload_numpy` because + # Sources are fresh per-call arrays, never rewritten by another + # ubatch. Their one-time uploads stay outside the persistent metadata + # registry. Via `upload_numpy` because # the indptrs are `T + 1` long: 64 KB at mnbt 16384, over the pageable # cliff at mnbt 131072. chunk_start_per_seq_gpu = upload_numpy(chunk_start_per_seq_np, device) @@ -4105,6 +4149,10 @@ def _build_paged_prefill_meta( attn_metadata.skip_prefix_len_csa = skip_csa_gpu attn_metadata.envelope_rows = envelope_rows + def _compress_publication_group(self, prefix=""): + groups = getattr(self.model_runner, "h2d_groups", None) + return None if groups is None else groups[f"{prefix}v4_plans"] + def _build_compress_plans( self, extend_lens_np, @@ -4114,6 +4162,7 @@ def _build_compress_plans( max_q_len: int | None = None, buf_prefix_ubatch: str = "", extra_write: int, + defer_to=None, ): """Build per-ratio CompressPlan dict consumed by batched compressor. @@ -4165,6 +4214,8 @@ def _build_compress_plans( context_lens_np, self._unique_compress_ratios_overlap, plan_buffers=plan_buffers, + publication_group=self._compress_publication_group(buf_prefix_ubatch), + defer_to=defer_to, running_bs=running_bs, max_q_len=max_q_len, extra_write=extra_write, @@ -4176,6 +4227,8 @@ def _populate_state_slot_mappings( scheduled_bs: int, running_bs: int, return_cpu: bool = False, + *, + publication_group=None, ): """Build the `[running_bs]` int32 tensor of per-request state slots. @@ -4210,7 +4263,15 @@ def _populate_state_slot_mappings( if len(pool_np) < scheduled_bs: pool_np = np.zeros(scheduled_bs, dtype=np.int32) slots_np = self._physical_slots(pool_np) - gpu = self._stage("v4_meta_state_slot_out", slots_np, pad_to=running_bs) + if publication_group is None: + gpu = self._stage("v4_meta_state_slot_out", slots_np, pad_to=running_bs) + else: + gpu = self._stage( + "v4_meta_state_slot_out", + slots_np, + pad_to=running_bs, + publication_group=publication_group, + ) if return_cpu: return gpu, slots_np return gpu @@ -4365,8 +4426,8 @@ def build_for_cudagraph_capture( cu_seqlens_q_gpu = var["cu_seqlens_q"].copy_to_gpu(bs + 1) var["context_lens"].np[:bs] = context_lens_np context_lens_gpu = var["context_lens"].copy_to_gpu(bs) - var["block_tables"].np[:bs] = block_tables_np - block_tables_gpu = var["block_tables"].copy_to_gpu(bs) + block_table_state(var["block_tables"]).prepare(block_tables_np) + block_tables_gpu = block_table_state(var["block_tables"]).publish(bs) state_slot_gpu = self._stage("v4_meta_state_slot_out", state_slot_np) # Read side captured from its own persistent buffer: replay-time # prepare_decode refills it, so a fork can change the values without @@ -4471,10 +4532,27 @@ def _state_slot_buffers(max_bs, device, *, prefix="", read_side=True): dtype=torch.int32, device=device, pin_memory=torch.device(device).type != "cpu", + publication_group=f"{prefix}v4_state", ) for direction in directions } + def _indexer_staging_buffers(self, bs, mnbt, *, prefix=""): + i32 = { + "dtype": torch.int32, + "device": self.device, + "publication_group": f"{prefix}v4_indexer", + } + buffers = {f"{prefix}v4_indexer_cu_committed": CpuGpuBuffer(bs + 1, **i32)} + if getattr(self, "indexer_layout", None) == FP4_GFX1250_NATURAL: + # Each nonempty query chunk contributes at most one extra request + # fragment per boundary and one leading zero. Across all chunks: + # entries <= bs + 2 * chunks (+ one PCP dummy request). + buffers[f"{prefix}v4_indexer_chunk_cu"] = CpuGpuBuffer( + bs + 2 * mnbt + 2, **i32 + ) + return buffers + def _alloc_v4_metadata_buffers(self) -> None: """Pre-allocate every buffer the V4 metadata builder writes into. @@ -4569,7 +4647,8 @@ def _alloc_v4_metadata_buffers(self) -> None: # what makes it CUDAGraph-safe (re-copied into the captured buffer # before graph.replay). The constant numpy source is precomputed once so # the per-fwd cost is a slice + H2D. - bufs["v4_qo_indptr"] = CpuGpuBuffer(T_dec + 1, **i32) + bufs["v4_qo_indptr"] = CpuGpuBuffer(T_dec + 1, publication_group="v4_qo", **i32) + bufs["v4_draft_qo_indptr"] = torch.arange(T_dec + 1, **i32) self._v4_qo_indptr_np = np.arange(T_dec + 1, dtype=np.int32) # Immutable, device-only empty CSR for reusing the H=128 sparse-prefill # ASM kernel in decode. Shared read-only across layers and TBO ubatches. @@ -4591,7 +4670,10 @@ def _alloc_v4_metadata_buffers(self) -> None: # int32 — `cp_gather_indexer_k_quant_cache` kernel signature is `int32_t*` # for cu_seq_lens. Also reused as cu_starts/cu_ends for fp8_mqa_logits # (which accepts both int32 and int64). - bufs["v4_indexer_cu_committed"] = CpuGpuBuffer(bs + 1, **i32) + bufs.update(self._indexer_staging_buffers(bs, mnbt)) + token_map = self.model_runner.forward_vars.get("batch_id_per_q_token") + if token_map is not None: + token_map.publication_group = "v4_tokens" # FP4 indexer decode plans use caller-held buffers: CUDAGraph captures # both their addresses and the launch grid. Keep one set per kernel # instance because qlen=1 decode uses a one-wave CTA while speculative @@ -4668,8 +4750,12 @@ def _alloc_v4_metadata_buffers(self) -> None: K_pool = (2 if is_overlap else 1) * ratio max_compress = mnbt // ratio + bs max_write = min(mnbt, bs * K_pool) - bufs[f"v4_compress_plan_{ratio}"] = CpuGpuBuffer(max_compress, 4, **i32) - bufs[f"v4_write_plan_{ratio}"] = CpuGpuBuffer(max_write, 4, **i32) + bufs[f"v4_compress_plan_{ratio}"] = CpuGpuBuffer( + max_compress, 4, publication_group="v4_plans", **i32 + ) + bufs[f"v4_write_plan_{ratio}"] = CpuGpuBuffer( + max_write, 4, publication_group="v4_plans", **i32 + ) # Pre-fill with sentinel so capture-time buffer state is valid # even before the first non-empty fwd. bufs[f"v4_compress_plan_{ratio}"].cpu.fill_(-1) @@ -4684,17 +4770,19 @@ def _alloc_v4_metadata_buffers(self) -> None: if getattr(self.model_runner.config, "enable_tbo", False) or getattr( self.model_runner.config, "enable_tbo_decode", False ): - self._alloc_v4_ubatch_decode_buffers(bufs, i32, i64) + self._alloc_v4_ubatch_metadata_buffers(bufs, i32, i64) self.model_runner.forward_vars.update(bufs) - def _alloc_v4_ubatch_decode_buffers(self, bufs: dict, i32: dict, i64: dict) -> None: - """Clone decode-path metadata buffers into ``ub{0,1}_`` prefixed sets. + def _alloc_v4_ubatch_metadata_buffers( + self, bufs: dict, i32: dict, i64: dict + ) -> None: + """Allocate metadata for TBO prefill/decode in ``ub{0,1}_`` sets. Mirrors the sizes chosen in :meth:`_alloc_v4_metadata_buffers` for the decode-relevant buffers plus the global per-fwd inputs the decode helpers read (``positions`` / ``context_lens`` / ``block_tables`` / - ``cu_seqlens_q``). Only invoked when ``enable_tbo_decode`` is set. + ``cu_seqlens_q``). Invoked for either prefill or decode TBO. """ mnbt = self.max_num_batched_tokens bs = self.max_bs @@ -4705,10 +4793,21 @@ def _alloc_v4_ubatch_decode_buffers(self, bufs: dict, i32: dict, i64: dict) -> N p = f"ub{ub_idx}_" # Global per-fwd decode inputs (live in model_runner.forward_vars # for the non-TBO path; cloned here so each ubatch slices its own). - bufs[f"{p}positions"] = CpuGpuBuffer(T_dec, **i64) - bufs[f"{p}context_lens"] = CpuGpuBuffer(bs, **i32) - bufs[f"{p}block_tables"] = CpuGpuBuffer(bs, self.block_table_cols, **i32) - bufs[f"{p}cu_seqlens_q"] = CpuGpuBuffer(bs + 1, **i32) + bufs[f"{p}positions"] = CpuGpuBuffer( + T_dec, publication_group=f"{p}v4_inputs", **i64 + ) + bufs[f"{p}context_lens"] = CpuGpuBuffer( + bs, publication_group=f"{p}v4_inputs", **i32 + ) + bufs[f"{p}block_tables"] = CpuGpuBuffer( + bs, self.block_table_cols, publication_group=f"{p}v4_inputs", **i32 + ) + bufs[f"{p}cu_seqlens_q"] = CpuGpuBuffer( + bs + 1, publication_group=f"{p}v4_inputs", **i32 + ) + bufs[f"{p}v4_qo_indptr"] = CpuGpuBuffer( + T_dec + 1, publication_group=f"{p}v4_inputs", **i32 + ) # V4 decode metadata buffers. bufs.update(self._state_slot_buffers(bs, self.device, prefix=p)) @@ -4729,8 +4828,10 @@ def _alloc_v4_ubatch_decode_buffers(self, bufs: dict, i32: dict, i64: dict) -> N ) for name in _DEST_ROW_BUFFERS.values(): bufs[f"{p}{name}"] = CpuGpuBuffer(T_dec, **i32) - bufs[f"{p}batch_id_per_q_token"] = CpuGpuBuffer(mnbt, **i32) - bufs[f"{p}v4_indexer_cu_committed"] = CpuGpuBuffer(bs + 1, **i32) + bufs[f"{p}batch_id_per_q_token"] = CpuGpuBuffer( + mnbt, publication_group=f"{p}v4_indexer", **i32 + ) + bufs.update(self._indexer_staging_buffers(bs, mnbt, prefix=p)) if self._indexer_fp4: if self.indexer_layout == FP4_GFX1250_NATURAL: from aiter.ops.opus.pa_mqa_logits_mxfp4 import ( @@ -4758,8 +4859,12 @@ def _alloc_v4_ubatch_decode_buffers(self, bufs: dict, i32: dict, i64: dict) -> N K_pool = (2 if is_overlap else 1) * ratio max_compress = mnbt // ratio + bs max_write = min(mnbt, bs * K_pool) - cbuf = CpuGpuBuffer(max_compress, 4, **i32) - wbuf = CpuGpuBuffer(max_write, 4, **i32) + cbuf = CpuGpuBuffer( + max_compress, 4, publication_group=f"{p}v4_plans", **i32 + ) + wbuf = CpuGpuBuffer( + max_write, 4, publication_group=f"{p}v4_plans", **i32 + ) cbuf.cpu.fill_(-1) cbuf.copy_to_gpu() wbuf.cpu.fill_(-1) @@ -4775,7 +4880,9 @@ def _dest_row_buffers(self, buf_prefix_ubatch: str = "") -> dict[int, torch.Tens for ratio, name in _DEST_ROW_BUFFERS.items() } - def _stage(self, name: str, arr, pad_to: int | None = None) -> torch.Tensor: + def _stage( + self, name: str, arr, pad_to: int | None = None, *, publication_group=None + ) -> torch.Tensor: """Write numpy `arr` into `forward_vars[name]` (CpuGpuBuffer) and return its GPU view sliced to len(arr). Asserts the buffer is large enough and that `arr.dtype` matches the buffer dtype (callers must @@ -4784,6 +4891,9 @@ def _stage(self, name: str, arr, pad_to: int | None = None) -> torch.Tensor: `pad_to` zero-fills the tail out to a wider view -- the padded batch a drafter runs. Zero is a real slot, so the caller owes those rows a reason they are never read. + + With an explicit `publication_group`, only fill its host source and + count. The enclosing builder publishes the group before any GPU use. """ buf = self.model_runner.forward_vars[name] n = arr.shape[0] if arr.ndim > 0 else 1 @@ -4802,8 +4912,13 @@ def _stage(self, name: str, arr, pad_to: int | None = None) -> torch.Tensor: f"V4 buffer {name!r} too small: need {width} padded, have {cap}. " f"Increase the corresponding bound in _alloc_v4_metadata_buffers." ) + if buf._publication is not None: + buf._publication.acquire_write() buf.np[n:width] = 0 buf.np[:n] = arr + if publication_group is not None: + publication_group.counts[publication_group.indices[name]] = width + return buf.gpu[:width] return buf.copy_to_gpu(width) @staticmethod diff --git a/atom/model_ops/attentions/gdn_attn.py b/atom/model_ops/attentions/gdn_attn.py index 7cf6cb49e7..e4c53b52c3 100644 --- a/atom/model_ops/attentions/gdn_attn.py +++ b/atom/model_ops/attentions/gdn_attn.py @@ -227,23 +227,10 @@ def _init_gdn_state( ), ) - self.spec_state_indices_tensor = CpuGpuBuffer( - (self.max_bs, self.num_spec + 1), - dtype=torch.int32, - device=self.device, - ) - self.non_spec_state_indices_tensor = CpuGpuBuffer( - (self.max_bs,), - dtype=torch.int32, - device=self.device, - ) - # Read side of a state fork. Only the prefill path can carry one (a - # fork is always followed by at least `min_fork_tokens` prompt tokens), - # so the spec/decode index buffers have no counterpart. - self.non_spec_state_indices_in_tensor = CpuGpuBuffer( - (self.max_bs,), - dtype=torch.int32, - device=self.device, + # All host state indices belong to the selected runner slot. The + # properties below follow PP rotation instead of retaining slot 0. + self.model_runner.forward_vars.update( + self._state_index_buffers(self.max_bs, self.num_spec, self.device) ) self.spec_sequence_masks = torch.ones( (self.max_bs,), @@ -280,8 +267,6 @@ def _init_gdn_state( ) gdn_metadata = { - "spec_state_indices": self.spec_state_indices_tensor, - "non_spec_state_indices": self.non_spec_state_indices_tensor, "spec_sequence_masks": self.spec_sequence_masks, "spec_token_indx": self.spec_token_indx, "non_spec_token_indx": self.non_spec_token_indx, @@ -291,6 +276,53 @@ def _init_gdn_state( } self.model_runner.forward_vars.update(gdn_metadata) + @staticmethod + def _state_index_buffers(max_bs, num_spec, device): + kwargs = { + "dtype": torch.int32, + "device": device, + "pin_memory": torch.device(device).type != "cpu", + "publication_group": "gdn_state", + } + return { + "spec_state_indices": CpuGpuBuffer(max_bs, num_spec + 1, **kwargs), + "non_spec_state_indices": CpuGpuBuffer(max_bs, **kwargs), + "non_spec_state_indices_in": CpuGpuBuffer(max_bs, **kwargs), + } + + @property + def spec_state_indices_tensor(self): + return self.model_runner.forward_vars["spec_state_indices"] + + @property + def non_spec_state_indices_tensor(self): + return self.model_runner.forward_vars["non_spec_state_indices"] + + @property + def non_spec_state_indices_in_tensor(self): + return self.model_runner.forward_vars["non_spec_state_indices_in"] + + def _check_state_indices_writable(self): + groups = getattr(self.model_runner, "h2d_groups", None) + if groups is not None: + groups["gdn_state"].check_writable() + + def _prepare_state_indices_for_capture(self, bs): + """Give Qwen4's synthetic requests distinct slots before capture.""" + self._check_state_indices_writable() + for indices in ( + self.non_spec_state_indices_tensor, + self.non_spec_state_indices_in_tensor, + ): + indices.np[:bs] = np.arange(bs, dtype=np.int32) + indices.copy_to_gpu(bs) + if self.use_spec_decode: + slots = self.spec_state_indices_tensor + slots.np[:bs] = np.arange(bs * slots.np.shape[1]).reshape(bs, -1) + if self.replayssm: + slots.np[:bs] = np.arange(bs)[:, None] + slots.copy_to_gpu(bs) + # ------------------------------------------------------------------ # # Per-request cache hooks (called from ModelRunner via base class). # # ------------------------------------------------------------------ # @@ -1060,6 +1092,7 @@ def prepare_state_indices(self, batch: ScheduledBatch, with_spec: bool = False): so this is where a contiguity assumption would have been *invented* rather than a place one has to be honoured. """ + self._check_state_indices_writable() non_spec_state_indices = self.non_spec_state_indices_tensor.np non_spec_state_indices_in = self.non_spec_state_indices_in_tensor.np spec_state_indices = self.spec_state_indices_tensor.np @@ -1146,18 +1179,18 @@ def prepare_gdn_metadata( else: self.prepare_state_indices(batch, with_spec=True) self.prepare_num_accepted_tokens(batch) + # The published query prefix (including graph padding) ends at + # the scheduled token count. Use its CPU value so ragged decode + # does not synchronize the device just to size an index view. spec_token_size = min( - num_decodes * (self.num_spec + 1), query_start_loc[-1].item() - ) - spec_token_indx = torch.arange( - spec_token_size, dtype=torch.int32, device=self.device - ) - non_spec_token_indx = torch.empty( - 0, dtype=torch.int32, device=query_start_loc.device - ) - spec_sequence_masks = torch.ones( - num_reqs, dtype=torch.bool, device=self.device + num_decodes * (self.num_spec + 1), batch.total_tokens_num ) + # The token range is immutable. Masks must be restored when a + # larger batch reuses entries padded False by a smaller batch. + spec_token_indx = self.spec_token_indx[:spec_token_size] + non_spec_token_indx = self.non_spec_token_indx[:0] + spec_sequence_masks = self.spec_sequence_masks[:num_reqs] + spec_sequence_masks.fill_(True) spec_state_indices_tensor = self.spec_state_indices_tensor.copy_to_gpu( num_reqs ) @@ -1294,31 +1327,21 @@ def _attach_gdn_decode_metadata( if self.use_spec_decode: self.spec_state_indices_tensor.gpu[num_decodes:, :].fill_(PAD_SLOT_ID) - self.spec_sequence_masks[:num_decodes].copy_( - gdn_metadata.spec_sequence_masks, non_blocking=True - ) self.spec_sequence_masks[num_decodes:].fill_(False) - gdn_metadata.spec_sequence_masks = self.spec_sequence_masks[:num_decodes] - - self.spec_token_indx[: gdn_metadata.spec_token_indx.size(0)].copy_( - gdn_metadata.spec_token_indx, non_blocking=True - ) - gdn_metadata.spec_token_indx = self.spec_token_indx[ - : gdn_metadata.spec_token_indx.size(0) - ] self.spec_query_start_loc[: num_decodes + 1].copy_( gdn_metadata.spec_query_start_loc[: num_decodes + 1], non_blocking=True ) - spec_num_query_tokens = self.spec_query_start_loc[num_decodes] - self.spec_query_start_loc[num_decodes + 1 :].fill_(spec_num_query_tokens) + # Broadcast the final prefix directly: fill_(GPU scalar) creates + # a temporary device allocation on ROCm. + self.spec_query_start_loc[num_decodes + 1 :].copy_( + self.spec_query_start_loc[num_decodes : num_decodes + 1], + non_blocking=True, + ) gdn_metadata.spec_query_start_loc = self.spec_query_start_loc[ : num_decodes + 1 ] - self.num_accepted_tokens[:num_decodes].copy_( - gdn_metadata.num_accepted_tokens[:num_decodes], non_blocking=True - ) self.num_accepted_tokens[num_decodes:].fill_(1) gdn_metadata.num_accepted_tokens = self.num_accepted_tokens[:num_decodes] else: @@ -1329,8 +1352,9 @@ def _attach_gdn_decode_metadata( gdn_metadata.non_spec_query_start_loc[: num_decodes + 1], non_blocking=True, ) - self.non_spec_query_start_loc[num_decodes + 1 :].fill_( - gdn_metadata.non_spec_query_start_loc[num_decodes] + self.non_spec_query_start_loc[num_decodes + 1 :].copy_( + self.non_spec_query_start_loc[num_decodes : num_decodes + 1], + non_blocking=True, ) gdn_metadata.non_spec_query_start_loc = self.non_spec_query_start_loc[ : num_decodes + 1 @@ -1504,15 +1528,16 @@ def prepare_decode( # type: ignore[override] attn_metadata, positions = super().prepare_decode( batch, running_bs, running_tokens, max_seqlen_q ) - self.model_runner.forward_vars["cu_seqlens_q"].cpu[ - running_bs: - ] = batch.total_tokens_num_decode - # we fill the attn_metadata cu_seqlens_q here since aiter attn won't calc it for decode - attn_metadata.cu_seqlens_q = self.model_runner.forward_vars[ - "cu_seqlens_q" - ].copy_to_gpu(running_bs + 1) - - self._attach_gdn_decode_metadata(batch, attn_metadata) + # publish_cu_seqlens_q already filled and uploaded the padded prefix + # before prepare_input_ids. Reuse it without rewriting a borrowed source. + attn_metadata.cu_seqlens_q = self.model_runner.forward_vars["cu_seqlens_q"].gpu[ + : running_bs + 1 + ] + + # The common decode builder already filled the block-table source. + self._attach_gdn_decode_metadata( + batch, attn_metadata, prepare_block_tables=False + ) return attn_metadata, positions def prepare_mtp_decode( @@ -1583,10 +1608,10 @@ def build_for_cudagraph_capture(self, bs: int): attn_metadata.context_lens, create=True ) - positions = var["positions"].copy_to_gpu(bs) # A capture runs a full synthetic batch, so nothing is padded and the # scheduled shape is the running one. capture_tokens = bs * int(var["max_qlen"]) + positions = var["positions"].copy_to_gpu(capture_tokens) context = Context( positions=positions, is_prefill=False, diff --git a/atom/model_ops/attentions/kimi_mla_gdn_attn.py b/atom/model_ops/attentions/kimi_mla_gdn_attn.py index f82f7ed48a..2a1bdc051d 100644 --- a/atom/model_ops/attentions/kimi_mla_gdn_attn.py +++ b/atom/model_ops/attentions/kimi_mla_gdn_attn.py @@ -1,7 +1,6 @@ # SPDX-License-Identifier: MIT # Copyright (C) 2024-2025, Advanced Micro Devices, Inc. All rights reserved. -import numpy as np import torch from aiter.dist.parallel_state import get_tp_group @@ -428,14 +427,12 @@ def prepare_decode( ) return attn_metadata, positions - def build_for_cudagraph_capture(self, bs: int): - if self.block_size == 1: - var = self.model_runner.forward_vars - var["kv_indptr"].np[: bs + 1] = np.arange(bs + 1, dtype=np.int32) - var["kv_indptr"].copy_to_gpu(bs + 1) - var["kv_indices"].gpu[:bs].zero_() - var["kv_last_page_lens"].gpu[:bs].fill_(1) + def _capture_needs_nonempty_kv(self, max_q_len: int) -> bool: + # Kimi's dense MLA warmup needs a page even with page_size=1. + # The parent owns the single upload, including DCP + MTP capture. + return True + def build_for_cudagraph_capture(self, bs: int): attn_metadata, context = super().build_for_cudagraph_capture(bs) attn_metadata.gdn_metadata = self._build_gdn_capture_metadata(bs) return attn_metadata, context diff --git a/atom/model_ops/attentions/qwen4_exp_attn.py b/atom/model_ops/attentions/qwen4_exp_attn.py index 09e13f474b..a583f0f2c1 100644 --- a/atom/model_ops/attentions/qwen4_exp_attn.py +++ b/atom/model_ops/attentions/qwen4_exp_attn.py @@ -35,6 +35,7 @@ qsa_draft_decode_metadata, ) from atom.utils import CpuGpuBuffer +from atom.utils.block_tables import block_table_state from .gdn_attn import GDNAttentionBackend, GDNAttentionMetadataBuilder from .pool_layout.entry_arena import EntryField, LayerMajorArena, entry_bytes_for @@ -190,10 +191,10 @@ def __init__(self, model_runner, **kwargs): i32 = {"dtype": torch.int32, "device": self.device} i64 = {"dtype": torch.int64, "device": self.device} self.model_runner.forward_vars["qsa_token_to_req"] = CpuGpuBuffer( - max_tokens, **i32 + max_tokens, publication_group="qsa", **i32 ) self.model_runner.forward_vars["qsa_logical_positions"] = CpuGpuBuffer( - max_tokens, **i64 + max_tokens, publication_group="qsa", **i64 ) # Written in place, never reallocated: a captured decode graph bakes in # the address it reads the compressed slots from. @@ -201,7 +202,7 @@ def __init__(self, model_runner, **kwargs): max_tokens, **i64 ) self.model_runner.forward_vars["ple_has_initial_state"] = CpuGpuBuffer( - self.max_bs, dtype=torch.bool, device=self.device + self.max_bs, dtype=torch.bool, device=self.device, publication_group="ple" ) # ------------------------------------------------------------------ # @@ -436,6 +437,7 @@ def _build_qsa_metadata( `mapped_tokens` is how many of `num_tokens` are real; the rest are CUDA-graph padding and are marked `-1` so no cache row is touched. """ + self._check_metadata_writable("qsa") token_to_req = self.model_runner.forward_vars["qsa_token_to_req"].np logical = self.model_runner.forward_vars["qsa_logical_positions"].np token_to_req[:num_tokens] = -1 @@ -523,6 +525,7 @@ def _build_ple_metadata( # A cold first chunk must not fold in whatever the recycled state slot # still held; anything with cached tokens continues its own window. + self._check_metadata_writable("ple") has_initial = self.model_runner.forward_vars["ple_has_initial_state"].np if is_prefill: has_initial[:num_reqs] = ( @@ -550,10 +553,10 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): # QSA reads the compressed cache through the block table on every # forward, so unlike dense prefill it cannot wait for `has_cached`. if attn_metadata.block_tables is None and batch.block_tables: - self.prepare_block_tables(batch) - attn_metadata.block_tables = self.model_runner.forward_vars[ - "block_tables" - ].copy_to_gpu(num_reqs) + # The parent already packed this source; only its upload was optional. + attn_metadata.block_tables = block_table_state( + self.model_runner.forward_vars["block_tables"] + ).publish(num_reqs) if attn_metadata.block_tables is None: attn_metadata.qsa_metadata = None attn_metadata.ple_metadata = None @@ -591,15 +594,6 @@ def prepare_decode( attn_metadata, positions = super().prepare_decode( batch, running_bs, running_tokens, max_seqlen_q ) - if positions.ndim == 2 and positions.shape[-1] < running_tokens: - # A captured mRoPE view uses the padded token count as its axis - # stride. Preserve that layout even when fewer rows are scheduled. - real_tokens = positions.shape[-1] - real_positions = self._mrope_cpu_view(real_tokens).copy() - padded_positions = self._mrope_cpu_view(running_tokens) - padded_positions.fill(0) - padded_positions[:, :real_tokens] = real_positions - positions = self._copy_mrope_to_gpu(running_tokens)[:, :real_tokens] bs = running_bs query_len = attn_metadata.max_seqlen_q num_tokens = running_tokens @@ -696,25 +690,14 @@ def build_for_cudagraph_capture(self, bs: int): # Capture has no scheduled state slots. Give its synthetic requests # distinct slots so the in-place windows cannot race on slot zero. # Normal metadata preparation overwrites these same buffers on replay. - for indices in ( - self.non_spec_state_indices_tensor, - self.non_spec_state_indices_in_tensor, - ): - indices.np[:bs] = np.arange(bs, dtype=np.int32) - indices.copy_to_gpu(bs) - if self.use_spec_decode: - slots = self.spec_state_indices_tensor - slots.np[:bs] = np.arange(bs * slots.np.shape[1]).reshape(bs, -1) - if self.replayssm: - slots.np[:bs] = np.arange(bs)[:, None] - slots.copy_to_gpu(bs) + self._prepare_state_indices_for_capture(bs) attn_metadata, context = super().build_for_cudagraph_capture(bs) runner = self.model_runner num_tokens = bs * int(attn_metadata.max_seqlen_q) attn_metadata.slot_mapping = runner.forward_vars["slot_mapping"].gpu[ :num_tokens ] - context.positions = runner.forward_vars["positions"].copy_to_gpu(num_tokens) + context.positions = runner.forward_vars["positions"].gpu[:num_tokens] attn_metadata.qsa_metadata = self._build_qsa_metadata( attn_metadata, @@ -731,6 +714,7 @@ def build_for_cudagraph_capture(self, bs: int): attn_metadata.ple_metadata = None return attn_metadata, context state_indices_in, state_indices_out = slots + self._check_metadata_writable("ple") self.model_runner.forward_vars["ple_has_initial_state"].np[:bs] = True attn_metadata.ple_metadata = Qwen4ExpPLEMetadata( query_start_loc=attn_metadata.cu_seqlens_q[: bs + 1], diff --git a/atom/model_ops/engram/device/hashing.py b/atom/model_ops/engram/device/hashing.py index e8ea8f2861..f940d6a3ee 100644 --- a/atom/model_ops/engram/device/hashing.py +++ b/atom/model_ops/engram/device/hashing.py @@ -123,8 +123,13 @@ def _engram_snapshot_kernel( tokens, padded_tokens, history_stride, + cursor_positions, + cursor_out, + cursor_slot_stride, + cursor_row_stride, NGRAM: tl.constexpr, BLOCK: tl.constexpr, + STAGE_CURSOR: tl.constexpr, ): token = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) live = token < tokens @@ -148,9 +153,29 @@ def _engram_snapshot_kernel( tl.where(live, source, -2), token < padded_tokens, ) + if STAGE_CURSOR and shift < NGRAM - 1: + # Snapshot is newest-first; a cursor stores oldest-first history + # AFTER this token. Both use these same committed lookbacks. + tl.store( + cursor_out + + batch * cursor_slot_stride + + local * cursor_row_stride + + NGRAM + - 1 + - shift, + source, + live, + ) + if STAGE_CURSOR: + position = tl.load(cursor_positions + token - local, live, 0) + tl.store( + cursor_out + batch * cursor_slot_stride + local * cursor_row_stride, + position + local + 1, + live, + ) -def engram_snapshot(tables, batch, out): +def engram_snapshot(tables, batch, out, *, cursor_positions=None, cursor_out=None): """Freeze lookbacks before the runner advances the committed cursor. The output has a stable address across graph buckets and replays. Only @@ -162,6 +187,22 @@ def engram_snapshot(tables, batch, out): raise ValueError("Engram snapshot must cover all tokens and lookbacks") if not out.is_contiguous(): raise ValueError("Engram snapshot must be contiguous") + stage_cursor = cursor_out is not None + if stage_cursor: + _check_plane(batch.history, tables, cursor_out) + if ( + cursor_out.ndim != 3 + or cursor_out.shape[0] < batch.history_index.numel() + or cursor_out.shape[-1] != tables.ngram + or cursor_positions is None + or cursor_positions.numel() < tokens + ): + raise ValueError("Engram snapshot requires per-request cursor candidates") + if ( + cursor_out.untyped_storage().data_ptr() + == batch.history.untyped_storage().data_ptr() + ): + raise ValueError("Engram snapshot cursor candidates must not alias history") if out.shape[0]: _engram_snapshot_kernel[(triton.cdiv(out.shape[0], 128),)]( batch.compressed, @@ -173,8 +214,13 @@ def engram_snapshot(tables, batch, out): tokens, out.shape[0], batch.history.stride(0), + cursor_positions, + cursor_out, + cursor_out.stride(0) if stage_cursor else 0, + cursor_out.stride(1) if stage_cursor else 0, NGRAM=tables.ngram, BLOCK=128, + STAGE_CURSOR=stage_cursor, ) return out diff --git a/atom/model_ops/engram/device/runtime.py b/atom/model_ops/engram/device/runtime.py index f9134f3456..7eb71648d3 100644 --- a/atom/model_ops/engram/device/runtime.py +++ b/atom/model_ops/engram/device/runtime.py @@ -221,6 +221,8 @@ def prepare( token_mask=None, padded_rows=None, batch=None, + cursor_positions=None, + cursor_out=None, ): """Stage one embedding row per row the forward will run. @@ -239,7 +241,16 @@ def prepare( compressed_rows = [] rows = token_ids.numel() if padded_rows is None else padded_rows if self.host.overlap is not None and (batch is not None or dummy): - return EngramInputs(self.host.overlap.prepare(batch, rows), histories, ()) + return EngramInputs( + self.host.overlap.prepare( + batch, + rows, + cursor_positions=cursor_positions, + cursor_out=cursor_out, + ), + histories, + (), + ) if dummy: self.host.stage_dummy(rows) next_histories = histories diff --git a/atom/model_ops/engram/device/staging.py b/atom/model_ops/engram/device/staging.py index 05c8734286..4d8d37e6fe 100644 --- a/atom/model_ops/engram/device/staging.py +++ b/atom/model_ops/engram/device/staging.py @@ -95,7 +95,7 @@ def _init_collective(self): "engram: hash/UVA/TP gather on one side stream with private IPC state" ) - def prepare(self, batch, width): + def prepare(self, batch, width, *, cursor_positions=None, cursor_out=None): if not 0 <= width <= self.host.max_num_tokens: raise ValueError("Engram staging exceeds capacity") snapshot = self.snapshot[:width] @@ -103,7 +103,13 @@ def prepare(self, batch, width): # Capture uses serving's kernels without accessing synthetic state. snapshot.fill_(-2) else: - engram_snapshot(self.uva.hash_tables, batch, snapshot) + engram_snapshot( + self.uva.hash_tables, + batch, + snapshot, + cursor_positions=cursor_positions, + cursor_out=cursor_out, + ) return EngramStagedRows(self, width) def start(self, width): diff --git a/atom/model_ops/sampler.py b/atom/model_ops/sampler.py index 2680a5de30..d07bc6e321 100644 --- a/atom/model_ops/sampler.py +++ b/atom/model_ops/sampler.py @@ -31,19 +31,20 @@ SAMPLER_EPS = 1e-10 -def _greedy_tokens(probs: torch.Tensor, greedy_mask: torch.Tensor) -> torch.Tensor: - """Argmax token for the rows `greedy_mask` selects, as int32. - - Reduces every row and keeps the selected answers. The obvious spelling, - `probs[greedy_mask].argmax(-1)`, gathers a `[greedy, vocab]` copy first, and - that copy costs more than reducing the rows it would have dropped: measured - at 256 rows of 200064 with three quarters greedy, 134.6us against 71.9. - - `tie="low"` is what makes this the pick `torch.argmax` made -- the default - promises no direction among equal scores, so a near-tie would resolve - differently from one TP rank to the next. +def _apply_greedy_tokens( + probs: torch.Tensor, temperatures: torch.Tensor, next_tokens: torch.Tensor +) -> torch.Tensor: + """Keep row selection on the device, with a fixed-size int32 result. + + A Python test of ``greedy_mask.any()`` waits for the GPU. Boolean indexing + also needs a data-dependent output size. Reduce all rows and select with + ``where`` so neither operation needs to read the mask back on the CPU. + ``tie="low"`` preserves argmax's lowest-token-ID tie breaking. """ - return topk_select(probs, 1, tie="low")[1].view(-1)[greedy_mask] + greedy_tokens = topk_select(probs, 1, tie="low")[1].view(-1) + return torch.where( + temperatures == 0, greedy_tokens, next_tokens.view(-1).to(torch.int) + ) def get_per_token_exponential(vocab_size: int, device) -> torch.Tensor: @@ -66,8 +67,8 @@ def sample_verification_tokens( logits: torch.Tensor, cu_num_draft_tokens: torch.Tensor, temperatures: torch.Tensor, - top_ks: torch.Tensor | None, - top_ps: torch.Tensor | None, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, ) -> torch.Tensor: """Sample each draft-conditioned target row with independent noise. @@ -85,7 +86,7 @@ def sample_verification_tokens( ) def expand(values): - if values is None or values.numel() == 1: + if not isinstance(values, torch.Tensor) or values.numel() == 1: return values return values[request_indices] @@ -102,8 +103,12 @@ def forward( self, logits: torch.Tensor, # (num_tokens, vocab_size) temperatures: torch.Tensor, # (num_tokens,) - top_ks: torch.Tensor | None = None, # (num_tokens,) int32, -1 means disabled - top_ps: torch.Tensor | None = None, # (num_tokens,) float32, 1.0 means disabled + top_ks: ( + int | torch.Tensor | None + ) = None, # (num_tokens,) int32, -1 means disabled + top_ps: ( + float | torch.Tensor | None + ) = None, # (num_tokens,) float32, 1.0 means disabled all_greedy: bool = False, # True if all temperatures are 0 (checked on CPU) needs_independent_noise: bool = False, ) -> torch.Tensor: # (num_tokens,) @@ -113,8 +118,8 @@ def forward( Args: logits: Raw logits from model (num_tokens, vocab_size) temperatures: Temperature for each token (num_tokens,), pre-clamped to eps - top_ks: Top-k value per token, -1 means disabled (num_tokens,) - top_ps: Top-p value per token, 1.0 means disabled (num_tokens,) + top_ks: Uniform CPU scalar or per-token tensor; -1 means disabled + top_ps: Uniform CPU scalar or per-token tensor; 1.0 means disabled all_greedy: True if all requests use greedy sampling (checked on CPU) needs_independent_noise: True when the batch contains fan-out siblings (SamplingParams.n>1). Forces fresh per-row random @@ -143,8 +148,8 @@ def forward( def _needs_filtering( self, - top_ks: torch.Tensor | None, - top_ps: torch.Tensor | None, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, ) -> bool: """Check if any request needs top-k or top-p filtering. @@ -187,8 +192,8 @@ def _topk_topp_sample( self, logits: torch.Tensor, temperatures: torch.Tensor, - top_ks: torch.Tensor | None, - top_ps: torch.Tensor | None, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, all_greedy: bool, needs_independent_noise: bool = False, ) -> torch.Tensor: @@ -228,23 +233,25 @@ def _topk_topp_sample( else: return self._native_sample(probs, top_ks, top_ps, temperatures) - def _to_tensor_scalar(self, x: torch.Tensor): - """Convert to (tensor, scalar) tuple for aiter ops. + def _to_tensor_scalar(self, x: float | torch.Tensor | None): + """Adapt filters to AITER's tensor/scalar arguments. - If tensor has size 1 (uniform value optimization from model_runner), - extract the scalar value for more efficient aiter kernel dispatch. + The runner supplies uniform filters as CPU scalars, avoiding a device + readback. Singleton tensors remain supported for existing callers. """ if x is None: return (None, 0) - if x.numel() == 1: # Uniform value - use scalar for efficiency - return (None, x[0].item()) + if not isinstance(x, torch.Tensor): + return (None, x) + if x.numel() == 1: + return (None, x.item()) return (x, 0) def _aiter_sample( self, probs: torch.Tensor, - top_ks: torch.Tensor, - top_ps: torch.Tensor, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, has_topk: bool, has_topp: bool, temperatures: torch.Tensor, @@ -278,22 +285,13 @@ def _aiter_sample( # Neither - just multinomial from probs next_tokens = torch.multinomial(probs, num_samples=1) - # Handle greedy sampling (temperature=0) - greedy_mask = temperatures == 0 - if greedy_mask.any(): - # Reduce the whole batch and keep the greedy rows, rather than - # `probs[greedy_mask]`, which first materializes a - # `[greedy, vocab]` copy -- the copy is most of what this used to - # cost, and the rows it drops are cheaper to reduce than to gather. - next_tokens[greedy_mask] = _greedy_tokens(probs, greedy_mask).unsqueeze(-1) - - return next_tokens.view(-1).to(torch.int) + return _apply_greedy_tokens(probs, temperatures, next_tokens) def _native_sample( self, probs: torch.Tensor, - top_ks: torch.Tensor, - top_ps: torch.Tensor, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, temperatures: torch.Tensor, ) -> torch.Tensor: """ @@ -324,15 +322,21 @@ def _native_sample( # The mask keeps tokens where cumsum - current_prob <= top_p # (i.e., before we exceed the threshold) if top_ps is not None: - topp_mask = (cumsum_probs - sorted_probs) <= top_ps.unsqueeze(-1) + p = top_ps.unsqueeze(-1) if isinstance(top_ps, torch.Tensor) else top_ps + topp_mask = (cumsum_probs - sorted_probs) <= p else: topp_mask = torch.ones_like(sorted_probs, dtype=torch.bool) # Top-k mask: keep first k tokens if top_ks is not None: indices = torch.arange(vocab_size, device=device).unsqueeze(0) - effective_k = torch.where(top_ks == -1, vocab_size, top_ks) - topk_mask = indices < effective_k.unsqueeze(-1) + if isinstance(top_ks, torch.Tensor): + effective_k = torch.where(top_ks == -1, vocab_size, top_ks).unsqueeze( + -1 + ) + else: + effective_k = vocab_size if top_ks == -1 else top_ks + topk_mask = indices < effective_k else: topk_mask = torch.ones_like(sorted_probs, dtype=torch.bool) @@ -349,12 +353,7 @@ def _native_sample( sampled_idx = torch.multinomial(filtered_probs, num_samples=1).squeeze(-1) next_tokens = sorted_indices.gather(1, sampled_idx.unsqueeze(-1)).squeeze(-1) - # Handle greedy (temperature=0) - greedy_mask = temperatures == 0 - if greedy_mask.any(): - next_tokens[greedy_mask] = _greedy_tokens(probs, greedy_mask) - - return next_tokens.to(torch.int) + return _apply_greedy_tokens(probs, temperatures, next_tokens) # Legacy methods kept for reference def greedy_sample( diff --git a/atom/model_ops/v4_kernels/compress_plan.py b/atom/model_ops/v4_kernels/compress_plan.py index c677b90c50..692ed50ee1 100644 --- a/atom/model_ops/v4_kernels/compress_plan.py +++ b/atom/model_ops/v4_kernels/compress_plan.py @@ -72,7 +72,7 @@ class CompressPlan: key_rope_positions_gpu: torch.Tensor | None = None # [≥num_compress] int64 or None -def _publish_key_rope_positions(plan_buffers, ratio, compress_buffer, count): +def _prepare_key_rope_positions(plan_buffers, ratio, compress_buffer, count, publish): """Where RoPE rotates each boundary's key, or None if nobody asked. Read back from the compress buffer AFTER its sentinel tail is filled, so @@ -84,7 +84,7 @@ def _publish_key_rope_positions(plan_buffers, ratio, compress_buffer, count): return None if count: buffer.np[:count] = (compress_buffer.np[:count, 2] // ratio) * ratio - return buffer.copy_to_gpu(count) + return publish(buffer, count) def plan_context_lens( @@ -109,6 +109,8 @@ def make_compress_plans( unique_ratios_overlap: Iterable[tuple[int, bool]], *, plan_buffers: dict, + publication_group=None, + defer_to=None, running_bs: int | None = None, max_q_len: int | None = None, decode_capacity_per_ratio: dict[int, int] | None = None, @@ -133,6 +135,12 @@ def make_compress_plans( calls (CUDAGraph requirement). Fresh per-call alloc is not supported — that pattern caused allocator-churn races (see `write_v4_paged_decode_indices` docstring). + publication_group: optional prebound owner group for these plan buffers. + All sources are filled before a single checked publication. + Independent callers without a group keep direct uploads. + defer_to: optional combined group that will publish these plan buffers. + The caller owns its counts and must publish before any GPU + plan consumer; used by V4.1's combined metadata preparation. running_bs: optional int — the CUDAGraph-padded batch size (>= bs). When PROVIDED this selects the DECODE CUDAGraph path: both `compress_plan_gpu` and `write_plan_gpu` are sliced to a @@ -187,6 +195,29 @@ def make_compress_plans( (fully sentinel-filled), so capture-time addresses match replay-time addresses even on a zero-token fwd. """ + # Check every source acquired by the owner before touching a pinned plan. A + # repeated call in the same forward must fail before overwriting an + # earlier publication, including when that call used checked direct. + if defer_to is not None and publication_group is None: + raise ValueError("Deferred plans require a publication group") + group = publication_group if defer_to is None else defer_to + if publication_group is not None: + publication_group.check_writable() + if defer_to is None: + for i in range(len(group.counts)): + group.counts[i] = None + + def publish(buffer, count): + if group is None: + return buffer.copy_to_gpu(count) + group.set_count(buffer, count) + return buffer.gpu[:count] + + def finish(): + if group is not None and defer_to is None: + group.publish(group.counts) + return out + bs = len(extend_lens_cpu) extend_lens_cpu = np.ascontiguousarray(extend_lens_cpu, dtype=np.int32) context_lens_cpu = np.ascontiguousarray(context_lens_cpu, dtype=np.int32) @@ -243,17 +274,17 @@ def _slices( if wcap > 0: wbuf.np[:wcap].fill(-1) out[ratio] = CompressPlan( - compress_plan_gpu=cbuf.copy_to_gpu(ccap), - write_plan_gpu=wbuf.copy_to_gpu(wcap), + compress_plan_gpu=publish(cbuf, ccap), + write_plan_gpu=publish(wbuf, wcap), num_compress=0, num_write=0, cu_compress_cpu=np.zeros(max(bs, 1) + 1, dtype=np.int32), compress_plan_cpu=None, - key_rope_positions_gpu=_publish_key_rope_positions( - plan_buffers, ratio, cbuf, ccap + key_rope_positions_gpu=_prepare_key_rope_positions( + plan_buffers, ratio, cbuf, ccap, publish ), ) - return out + return finish() # Per-token columns shared across ratios. batch_ids = np.repeat(np.arange(bs, dtype=np.int32), extend_lens_cpu) @@ -341,8 +372,8 @@ def _slices( wbuf.np[:n_write] = write_plan if write_slice > n_write: wbuf.np[n_write:write_slice].fill(-1) # sentinel - compress_plan_gpu = cbuf.copy_to_gpu(compress_slice) - write_plan_gpu = wbuf.copy_to_gpu(write_slice) + compress_plan_gpu = publish(cbuf, compress_slice) + write_plan_gpu = publish(wbuf, write_slice) out[ratio] = CompressPlan( compress_plan_gpu=compress_plan_gpu, @@ -351,8 +382,8 @@ def _slices( num_write=n_write, cu_compress_cpu=cu_compress, compress_plan_cpu=compress_plan if n_compress > 0 else None, - key_rope_positions_gpu=_publish_key_rope_positions( - plan_buffers, ratio, cbuf, compress_slice + key_rope_positions_gpu=_prepare_key_rope_positions( + plan_buffers, ratio, cbuf, compress_slice, publish ), ) - return out + return finish() diff --git a/atom/rollout/model_runner_ext.py b/atom/rollout/model_runner_ext.py index ef90222a31..dfd7ec58a4 100644 --- a/atom/rollout/model_runner_ext.py +++ b/atom/rollout/model_runner_ext.py @@ -86,8 +86,8 @@ def postprocess( batch: ScheduledBatch, logits: torch.Tensor, temperatures: torch.Tensor, - top_ks: torch.Tensor | None, - top_ps: torch.Tensor | None, + top_ks: int | torch.Tensor | None, + top_ps: float | torch.Tensor | None, all_greedy: bool, hidden_states: torch.Tensor, needs_independent_noise: bool = False, diff --git a/atom/spec_decode/drafter.py b/atom/spec_decode/drafter.py index 72bd2705b4..e313e27356 100644 --- a/atom/spec_decode/drafter.py +++ b/atom/spec_decode/drafter.py @@ -20,6 +20,7 @@ get_forward_context, set_forward_context, ) +from atom.utils.h2d import h2d_producer logger = logging.getLogger("atom") @@ -174,15 +175,34 @@ def __init__(self, atom_config: Config, device: torch.device, runner): self._captures_aux = False self._aux_buffers: list[torch.Tensor] = [] - i32_kwargs = {"dtype": torch.int32, "device": self.device} - i64_kwargs = {"dtype": torch.int64, "device": self.device} - max_bs = self.config.max_num_seqs - self.cu_num_draft_tokens = CpuGpuBuffer(max_bs, **i32_kwargs) - self.target_logits_indices = CpuGpuBuffer(max_bs * self.mtp_k, **i64_kwargs) - self.bonus_logits_indices = CpuGpuBuffer(max_bs, **i64_kwargs) + self.metadata_buffers = self._allocate_metadata_buffers( + self.config.max_num_seqs, self.mtp_k, self.device + ) self._build_draft_graphs() + @staticmethod + def _allocate_metadata_buffers(max_bs, mtp_k, device): + kwargs = { + "device": device, + "publication_group": "spec_decode", + "pin_memory": torch.device(device).type != "cpu", + } + return { + "cu_num_draft_tokens": CpuGpuBuffer(max_bs, dtype=torch.int32, **kwargs), + "target_logits_indices": CpuGpuBuffer( + max_bs * mtp_k, dtype=torch.int64, **kwargs + ), + "bonus_logits_indices": CpuGpuBuffer(max_bs, dtype=torch.int64, **kwargs), + # Device-only scratch, shared by PP slots just like other device + # tensors in forward_vars. The rejection sampler consumes it on + # the forward stream before the next prepare can overwrite it; + # draft proposal uses separate token storage. No host publication. + "verification_draft_token_ids": torch.empty( + max_bs * mtp_k, dtype=torch.int32, device=device + ), + } + # ---- draft passes ---- def _declare_draft_graphs(self) -> tuple[DraftGraph, ...]: """Declare this drafter's warmable forwards. Opt-in. @@ -438,6 +458,8 @@ def anchors_to_gpu(self, anchors: list[int]) -> torch.Tensor: """ n = len(anchors) buf = self.runner.forward_vars["draft_next_tokens"] + if buf._publication is not None: + buf._publication.acquire_write() buf.np[:n] = anchors return buf.copy_to_gpu(n) @@ -651,12 +673,18 @@ def prepare_inputs( return token_indices - def calc_spec_decode_metadata( + @h2d_producer("spec_decode", runner="runner") + def prepare_spec_decode_indices( self, num_sampled_tokens: np.ndarray, cu_num_sampled_tokens: np.ndarray, - input_ids: torch.Tensor, - ) -> SpecDecodeMetadata: + publication_group, + ) -> tuple[np.ndarray, int]: + """Fill host indices/counts before their group's first GPU consumer. + + This uses only the settled query lengths, so it can share token input + publication. The caller must publish the group before gathering IDs. + """ scheduled_bs = len(num_sampled_tokens) # num_draft = num_sampled - 1 per request. num_sampled_tokens is the @@ -685,19 +713,42 @@ def calc_spec_decode_metadata( # [0, 1, 2, 5, 6, 9] target_logits_indices += arange - # Do the CPU -> GPU copy. - self.target_logits_indices.np[:sum_drafted_tokens] = target_logits_indices - self.cu_num_draft_tokens.np[:scheduled_bs] = cu_num_draft_tokens - self.bonus_logits_indices.np[:scheduled_bs] = bonus_logits_indices - target_logits_indices = self.target_logits_indices.copy_to_gpu( - sum_drafted_tokens - ) - cu_num_draft_tokens = self.cu_num_draft_tokens.copy_to_gpu(scheduled_bs) - bonus_logits_indices = self.bonus_logits_indices.copy_to_gpu(scheduled_bs) + var = self.runner.forward_vars + var["target_logits_indices"].np[:sum_drafted_tokens] = target_logits_indices + var["cu_num_draft_tokens"].np[:scheduled_bs] = cu_num_draft_tokens + var["bonus_logits_indices"].np[:scheduled_bs] = bonus_logits_indices + group = publication_group + counts = group.counts + counts[group.indices["target_logits_indices"]] = sum_drafted_tokens + counts[group.indices["cu_num_draft_tokens"]] = scheduled_bs + counts[group.indices["bonus_logits_indices"]] = scheduled_bs + return num_draft_tokens, sum_drafted_tokens + + def calc_spec_decode_metadata( + self, + num_sampled_tokens: np.ndarray, + cu_num_sampled_tokens: np.ndarray, + input_ids: torch.Tensor, + *, + prepared_indices: tuple[np.ndarray, int] | None = None, + ) -> SpecDecodeMetadata: + if prepared_indices is None: + group = self.runner.h2d_groups["spec_decode"] + prepared_indices = self.prepare_spec_decode_indices( + num_sampled_tokens, cu_num_sampled_tokens, group + ) + group.publish(group.counts) + num_draft_tokens, sum_drafted_tokens = prepared_indices + scheduled_bs = len(num_draft_tokens) + var = self.runner.forward_vars + target_logits_indices = var["target_logits_indices"].gpu[:sum_drafted_tokens] + cu_num_draft_tokens = var["cu_num_draft_tokens"].gpu[:scheduled_bs] + bonus_logits_indices = var["bonus_logits_indices"].gpu[:scheduled_bs] # Compute the draft token ids. # draft_token_indices: [ 1, 2, 3, 105, 106, 208] - draft_token_ids = torch.index_select(input_ids[1:], 0, target_logits_indices) + draft_token_ids = var["verification_draft_token_ids"][:sum_drafted_tokens] + torch.index_select(input_ids[1:], 0, target_logits_indices, out=draft_token_ids) metadata = SpecDecodeMetadata( draft_token_ids=draft_token_ids, diff --git a/atom/utils/__init__.py b/atom/utils/__init__.py index 9b78086f75..18591189d0 100644 --- a/atom/utils/__init__.py +++ b/atom/utils/__init__.py @@ -697,6 +697,18 @@ def pack_rows(dst: np.ndarray, rows: Sequence) -> None: class CpuGpuBuffer: """Buffer to easily copy tensors between CPU and GPU.""" + def __setattr__(self, name, value): + if ( + name in ("cpu", "gpu", "np") + and self.__dict__.get("_publication") is not None + ): + from atom.utils.h2d import PublicationError + + raise PublicationError( + "bound metadata storage is fixed; recreate the buffer and owner" + ) + object.__setattr__(self, name, value) + def __init__( self, *size: int | torch.SymInt, @@ -704,7 +716,12 @@ def __init__( device: torch.device, pin_memory: bool = True, with_numpy: bool = True, + publication_group: str | None = None, + publication_unit: str = "rows", ) -> None: + self._publication = None + self.publication_group = publication_group + self.publication_unit = publication_unit self.cpu = torch.zeros(*size, dtype=dtype, device="cpu", pin_memory=pin_memory) self.gpu = torch.zeros_like(self.cpu, device=device) self.np: np.ndarray @@ -719,7 +736,11 @@ def __init__( ) self.np = self.cpu.numpy() - def copy_to_gpu(self, n: int | None = None) -> torch.Tensor: + def copy_to_gpu( + self, n: int | None = None, *, republish_reason: str | None = None + ) -> torch.Tensor: + if self._publication is not None: + return self._publication.copy_to_gpu(n, republish_reason=republish_reason) if n is None: return self.gpu.copy_(self.cpu, non_blocking=True) return self.gpu[:n].copy_(self.cpu[:n], non_blocking=True) @@ -742,6 +763,8 @@ def clone(self) -> "CpuGpuBuffer": device=self.gpu.device, pin_memory=self.cpu.is_pinned(), with_numpy=hasattr(self, "np"), + publication_group=self.publication_group, + publication_unit=self.publication_unit, ) new.cpu.copy_(self.cpu) new.gpu.copy_(self.gpu) diff --git a/atom/utils/block_tables.py b/atom/utils/block_tables.py new file mode 100644 index 0000000000..60c0fabc97 --- /dev/null +++ b/atom/utils/block_tables.py @@ -0,0 +1,233 @@ +# SPDX-License-Identifier: MIT +"""A forward slot's CPU page-table snapshot and its published GPU revision.""" + +import numpy as np + +from atom.model_engine.sequence import BlockTable + + +def _int32_row(row): + values = np.asarray(row) + if values.ndim != 1: + raise ValueError("Block-table rows must be one-dimensional") + if values.dtype != np.int32: + if values.size and ( + values.dtype.kind not in "iu" + or values.min() < 0 + or values.max() > np.iinfo(np.int32).max + ): + raise ValueError("PAGE ids must be nonnegative int32 integers") + values = values.astype(np.int32) + return np.ascontiguousarray(values) + + +class BlockTableState: + """Keep row versions and lengths beside the pinned snapshot, not its payload. + + All writers, including capture and padding, must use this object. Versions + and lengths identify rows without payload snapshots or per-row key tuples. + """ + + def __init__(self, buffer): + self.buffer = buffer + self.cpu = buffer.np + if ( + self.cpu.dtype != np.int32 + or self.cpu.ndim != 2 + or not self.cpu.flags.c_contiguous + ): + raise TypeError( + "Block tables require a contiguous two-dimensional int32 buffer" + ) + self.bytes = self.cpu.view(np.uint8) + self.flat = memoryview(self.cpu).cast("B").cast("i") if self.cpu.size else () + self.versions = [] + self.lengths = [] + self.zero_to = 0 + self.revision = 0 + self.published = None + self.page_limit = None + self.checked = None + + def _acquire(self): + binding = getattr(self.buffer, "_publication", None) + if binding is not None: + binding.acquire_write() + + def _begin_write(self): + self._acquire() + self.revision += 1 + self.published = None + # A failed host write must not leave reusable CPU keys or padding. + self.versions = [] + self.lengths = [] + self.zero_to = 0 + + def _normalize(self, rows, versions, lengths): + rows = list(rows) + for i, (row, version, length) in enumerate(zip(rows, versions, lengths)): + if version is not None: + continue + rows[i] = values = _int32_row(row) + unchanged = ( + i < len(self.lengths) + and length == self.lengths[i] + and np.array_equal(values, self.cpu[i, :length]) + ) + versions[i] = self.versions[i] if unchanged else object() + return rows + + def _validate_range(self, rows, versions, lengths, page_limit): + # Validation belongs to an append lineage, regardless of its row index. + # Reordering requests must not re-read their already checked pages. + checked = self.checked if page_limit == self.page_limit else {} + limit = min(page_limit, 1 << 31) + for row, version, length in zip(rows, versions, lengths): + start = checked.get(version, 0) + if length <= start: + continue + if length - start == 1: + invalid = not 0 <= row[start] < page_limit + else: + tail = _int32_row(row)[start:] + # Signed negatives sort after every valid PAGE id as uint32, + # so one reduction checks both bounds without a payload copy. + invalid = tail.view(np.uint32).max() >= limit + if invalid: + raise ValueError("Request PAGE table is incomplete or out of range") + return dict(zip(versions, lengths)) + + def prepare(self, rows, *, pad_to=None, page_limit=None, _keys=None): + n = len(rows) + capacity, columns = self.cpu.shape + end = n if pad_to is None else pad_to + if not 0 <= n <= end <= capacity: + raise ValueError("Block-table rows exceed the declared capacity") + if _keys is None: + versions = [ + row.version if isinstance(row, BlockTable) else None for row in rows + ] + lengths = list(map(len, rows)) + else: + versions, lengths = _keys + if ( + versions == self.versions + and lengths == self.lengths + and end <= self.zero_to + and page_limit == self.page_limit + ): + return self + if max(lengths, default=0) > columns: + raise ValueError("Block-table row exceeds the destination columns") + if None in versions: + rows = self._normalize(rows, versions, lengths) + checked = ( + self._validate_range(rows, versions, lengths, page_limit) + if page_limit is not None + else None + ) + + old_versions, old_lengths = self.versions, self.lengths + old_n = len(old_lengths) + if n > old_n: + old_versions = old_versions + [None] * (n - old_n) + old_lengths = old_lengths + [columns] * (n - old_n) + changes = [ + (i, length, old if version == previous and length >= old else 0, old) + for i, (version, previous, length, old) in enumerate( + zip(versions, old_versions, lengths, old_lengths) + ) + if version != previous or length != old + ] + clear_tail = end > n and (n != old_n or end > self.zero_to) + zero_to = max(end, self.zero_to) if n == old_n else end + if changes or clear_tail: + self._begin_write() # all validation precedes the first write + # A fully replaced batch can clear exposed tails in one pass. + # Partial updates retain every unaffected row and copied prefix. + bulk_clear = ( + len(changes) == n + and lengths != old_lengths + and min(lengths, default=columns) < columns + and not any(start for _, _, start, _ in changes) + ) + if bulk_clear: + self.bytes[:n] = 0 + self._copy_changes(rows, changes, columns, bulk_clear) + if clear_tail: + self.bytes[n:end] = 0 + self.versions, self.lengths = versions, lengths + self.zero_to = zero_to + self.page_limit = page_limit + self.checked = checked + return self + + def _copy_changes(self, rows, changes, columns, bulk_clear): + flat = self.flat + for i, length, start, old in changes: + if not bulk_clear and length < old: + self.bytes[i, length * 4 : old * 4] = 0 + base = i * columns + if length - start == 1: + flat[base + start] = rows[i][start] + elif length > start: + flat[base + start : base + length] = ( + rows[i][start:] if start else rows[i] + ) + + def pad(self, scheduled_bs, running_bs): + """Finalize the padded range before publication, including draft rows.""" + if not 0 <= scheduled_bs <= running_bs <= self.cpu.shape[0]: + raise ValueError("Invalid block-table padding range") + if scheduled_bs != len(self.lengths) or running_bs > self.zero_to: + versions = self.versions[:scheduled_bs] + lengths = self.lengths[:scheduled_bs] + self._begin_write() + self.bytes[scheduled_bs:running_bs] = 0 + self.versions, self.lengths = versions, lengths + self.zero_to = running_bs + return self + + def slice_to(self, buffer, start, count, *, pad_to): + """Prepare a TBO destination using the source snapshot's row revisions.""" + stop = start + count + if stop > len(self.lengths): + # Independent callers can supply an already packed CPU table. + return block_table_state(buffer).prepare( + self.cpu[start:stop], pad_to=pad_to + ) + rows = [self.cpu[i, : self.lengths[i]] for i in range(start, stop)] + keys = self.versions[start:stop], self.lengths[start:stop] + return block_table_state(buffer).prepare(rows, pad_to=pad_to, _keys=keys) + + def publish(self, count, *, group=None): + """Publish a changed table, optionally with the caller's metadata group. + + Omitting a table from a group does not omit its GPU view from metadata. + Record the revision only after the entire publication succeeds. + """ + if count is None: + if group is not None: + group.publish(group.counts) + return None + key = (self.revision, count) + previous = self.published + dirty = ( + previous is None or previous[0] is not self.buffer.gpu or previous[1] != key + ) + if group is not None: + group.set_count(self.buffer, count if dirty else None) + group.publish(group.counts) + elif dirty: + self.buffer.copy_to_gpu(count) + if dirty: + self.published = (self.buffer.gpu, key) + return self.buffer.gpu[:count] + + +def block_table_state(buffer): + """Get the state belonging to this physical forward buffer/PP slot.""" + state = getattr(buffer, "_block_table", None) + if state is None or state.cpu is not buffer.np: + state = buffer._block_table = BlockTableState(buffer) + return state diff --git a/atom/utils/envs.py b/atom/utils/envs.py index e371570f60..68d98aec82 100644 --- a/atom/utils/envs.py +++ b/atom/utils/envs.py @@ -46,6 +46,9 @@ def _positive_float_env(name: str, default: str) -> float: environment_variables: dict[str, Callable[[], Any]] = { + # Forward metadata transport: direct or packed. Both keep source checks. + # Single-member groups and strided bindings retain direct copies. + "ATOM_H2D_BACKEND": lambda: os.getenv("ATOM_H2D_BACKEND", "direct"), # Opt-in single-HCA engine pool: "auto" or explicit comma-separated HCAs. "ATOM_MOONCAKE_MATCHED_RAILS": lambda: os.getenv("ATOM_MOONCAKE_MATCHED_RAILS", ""), # Protect reused KV prefixes from one-off prefill scans. Opt-in. diff --git a/atom/utils/forward_context.py b/atom/utils/forward_context.py index cf50e1594c..577ff4472a 100644 --- a/atom/utils/forward_context.py +++ b/atom/utils/forward_context.py @@ -423,8 +423,9 @@ def _rows(name): t = getattr(attn_metadata, name, None) return None if t is None else int(t.shape[0]) - # `input_ids` is the argument, this rank's own rows -- the cudagraph - # branch re-slices the buffer to `running_tokens` itself. + # `input_ids` is the scheduled prefix at the runner boundary. Uniform + # decode exposes `running_tokens` to the model in both eager and graph + # execution; a mixed prefill/decode step keeps its local token count. assert input_ids.shape[0] == self.scheduled_tokens, ( f"input_ids length {input_ids.shape[0]} != scheduled_tokens=" f"{self.scheduled_tokens} ({self})" @@ -480,8 +481,7 @@ class Context: is_prefill: bool = False is_dummy_run: bool = False # What this rank was handed. Duplicated from `forward_mode` because a - # capture context has none; `scheduled_tokens` is what an eager forward - # actually runs. + # capture context has none; `scheduled_tokens` counts this rank's real rows. scheduled_bs: int = 0 scheduled_tokens: int = 0 # The step's DP-unified padded shape. `running_bs` counts SEQUENCES (graph diff --git a/atom/utils/h2d.py b/atom/utils/h2d.py new file mode 100644 index 0000000000..8f99f71563 --- /dev/null +++ b/atom/utils/h2d.py @@ -0,0 +1,486 @@ +# SPDX-License-Identifier: MIT +"""Checked publication of persistent host buffers on an owner's stream. + +The owner supplies its existing reuse event; this module owns no ring. CPU +writes must follow begin() or acquire_write(), before publish(). Tensor aliases +can still write around this protocol: producers must migrate with their buffers. +""" + +from __future__ import annotations + +import operator +import threading +from functools import wraps + +import torch + + +class PublicationError(RuntimeError): + """A publication violates ownership or asynchronous source lifetime.""" + + +def h2d_producer(*group_names, runner=None): + """Declare source groups a method writes, checking before its body runs. + + ``runner`` names the attribute holding the runner, or None for its own + methods. Resolve the current group on every call so PP slot rotation is + respected. Standalone producers without a publication registry retain + their unregistered behavior; a registered runner must provide the group. + This neither publishes data nor grants permission to republish it. + """ + get_runner = operator.attrgetter(runner) if runner else lambda instance: instance + + def decorate(produce): + @wraps(produce) + def checked(instance, *args, **kwargs): + groups = getattr(get_runner(instance), "h2d_groups", None) + if groups is not None: + for name in group_names: + groups[name].check_writable() + return produce(instance, *args, **kwargs) + + return checked + + return decorate + + +def _count(value, capacity, name): + if isinstance(value, bool): + raise TypeError(f"{name}: count must be an integer, not bool") + try: + value = operator.index(value) + except TypeError as exc: + raise TypeError(f"{name}: count must be an integer") from exc + if not 0 <= value <= capacity: + raise ValueError(f"{name}: count {value} is outside [0, {capacity}]") + return value + + +def _range(tensor): + # Reserve the whole span for strided views; overlapping holes cannot be + # registered as independent buffers. Direct still supports their layout. + start = tensor.data_ptr() + if not tensor.numel(): + return start, start + span = 1 + sum((n - 1) * s for n, s in zip(tensor.shape, tensor.stride())) + return start, start + span * tensor.element_size() + + +def _overlaps(a, b): + return a[0] < b[1] and b[0] < a[1] + + +def _reason(reason): + return isinstance(reason, str) and bool(reason.strip()) + + +class PublicationRegistry: + """Registration-time alias checks shared by a runner's owners and slots.""" + + def __init__(self): + self._bindings = [] + + def add(self, binding): + for prior in self._bindings: + if prior.device == binding.device and _overlaps( + prior.destination_range, binding.destination_range + ): + raise ValueError( + f"{binding.name}: destination overlaps registered {prior.name}" + ) + if _overlaps(prior.source_range, binding.source_range): + raise ValueError( + f"{binding.name}: source overlaps registered {prior.name}" + ) + self._bindings.append(binding) + + +class PublicationOwner: + """One existing owner/slot, with an epoch shared by all its publish groups. + + finish() seals host preparation, retaining the ledger until the next begin. + A later host phase must resume the SAME epoch, preserving duplicate checks. + begin() starts a forward, initialization or capture-preparation epoch; + an auxiliary group in the same forward must not reset it. + """ + + def __init__(self, device, completion=None, *, registry=None): + self.device = torch.device(device) + if self.device.type == "cuda" and self.device.index is None: + self.device = torch.device("cuda", torch.cuda.current_device()) + if self.device.type == "cuda" and completion is None: + raise ValueError("GPU publication requires the owner's reuse event") + self.completion = completion + self.registry = registry if registry is not None else PublicationRegistry() + self.epoch = 0 + self._state = "idle" + self._stream = None + self._thread = None + self._stream_id = None + # Compare PyTorch stream IDs without constructing a Stream wrapper. + # Native handles can alias distinct PyTorch stream identities. + self._get_current_stream = getattr(torch._C, "_cuda_getCurrentStream", None) + self._bindings = [] + self._groups = [] + + def bind(self, buffer, name, *, unit="rows"): + if self._state != "idle": + raise PublicationError("bind buffers before starting publication") + if getattr(buffer, "_publication", None) is not None: + raise ValueError(f"{name}: buffer is already bound") + if any(b.name == name for b in self._bindings): + raise ValueError(f"duplicate logical name: {name}") + binding = BufferPublication(self, buffer, name, unit) + self.registry.add(binding) + self._bindings.append(binding) + buffer._publication = binding + binding._single = self.group(name, (binding,)) + return binding + + def group(self, name, members): + if self._state != "idle": + raise PublicationError("create groups during initialization") + members = tuple(members) + if any(b.owner is not self for b in members): + raise ValueError("a group must belong to one owner/device") + if len({id(b) for b in members}) != len(members): + raise ValueError("a group cannot publish a destination twice") + group = PublicationGroup(self, name, members) + self._groups.append(group) + return group + + def use_packed_transport(self): + """Configure maximal eligible groups once; covered groups stay direct. + + The runner defines consumer boundaries. This layer chooses storage, + avoiding duplicate arenas for producer groups covered by a boundary. + Standalone calls to those groups still use checked direct copies. + """ + if self._state != "idle": + raise PublicationError("choose the transport during initialization") + packed_members = [] + configured = [] + for group in sorted(self._groups, key=lambda g: -len(g.members)): + members = set(group.members) + if len(members) < 2 or any(members <= packed for packed in packed_members): + continue + group.use_transport("packed") + if group.transport == "packed": + packed_members.append(members) + configured.append(group) + return configured + + def _current_stream(self): + if self.device.type != "cuda": + return None + if torch.cuda.current_device() != self.device.index: + raise PublicationError("publication on the wrong device") + if torch.cuda.is_current_stream_capturing(): + raise PublicationError("publish metadata before actual graph capture") + return torch.cuda.current_stream(self.device) + + def _check_active(self): + if self._state != "active": + raise PublicationError(f"owner is {self._state}; acquire sources first") + if threading.get_ident() != self._thread: + raise PublicationError("publish on the owner thread") + if self.device.type == "cuda": + if torch.cuda.current_device() != self.device.index: + raise PublicationError("publication on the wrong device") + if torch.cuda.is_current_stream_capturing(): + raise PublicationError("publish metadata before actual graph capture") + if self._get_current_stream is not None: + same_stream = ( + self._get_current_stream(self.device.index)[0] == self._stream_id + ) + else: + same_stream = torch.cuda.current_stream(self.device) == self._stream + if not same_stream: + raise PublicationError("publish on the owner's current compute stream") + + def begin(self): + """Wait for source reuse BEFORE producer writes; start a new epoch.""" + if self._state not in ("idle", "sealed"): + raise PublicationError(f"cannot begin: owner is {self._state}") + stream = self._current_stream() + if self.completion is not None: + # The caller records constructor uploads/slot clones before first use. + try: + self.completion.synchronize() + except BaseException: + self._state = "failed" + raise + self.epoch += 1 + self._stream = stream + self._stream_id = None if stream is None else stream.stream_id + self._thread = threading.get_ident() + for binding in self._bindings: + binding._writable = True + for group in self._groups: + group._header_busy = False + self._state = "active" + + def finish(self): + """Seal preparation and cover every queued source read.""" + self._check_active() + try: + if self.completion is not None: + self.completion.record(self._stream) + except BaseException: + self._state = "failed" + raise + self._state = "sealed" + + def resume(self): + """Reopen a later phase without resetting the forward's ledger.""" + if self._state != "sealed": + raise PublicationError(f"cannot resume: owner is {self._state}") + if threading.get_ident() != self._thread: + raise PublicationError("resume on the original owner thread") + if self._current_stream() != self._stream: + raise PublicationError("resume on the original owner stream") + self._state = "active" + # Sources stay protected; acquire_write must precede their next write. + + def _wait_sources(self): + self._check_active() + try: + if self.completion is not None: + self.completion.record(self._stream) + self.completion.synchronize() + except BaseException: + self._state = "failed" + raise + for group in self._groups: + group._header_busy = False + + def fail(self): + """Retain storage and reject reuse after partially submitted failures.""" + self._state = "failed" + + def drain(self): + """Drain failed work; the owner remains failed and must be rebuilt.""" + if self._state != "failed": + raise PublicationError("drain is only for failed publication owners") + if self._stream is not None: + self._stream.synchronize() + + +class BufferPublication: + def __init__(self, owner, buffer, name, unit): + self.owner, self.buffer, self.name = owner, buffer, name + self.source, self.destination = buffer.cpu, buffer.gpu + self.device = self.destination.device + if self.device != owner.device or self.source.device.type != "cpu": + raise ValueError(f"{name}: incorrect source/destination device") + if ( + self.source.shape != self.destination.shape + or self.source.dtype != self.destination.dtype + ): + raise ValueError(f"{name}: publication cannot reshape or cast") + if unit not in ("rows", "elements", "bytes"): + raise ValueError(f"{name}: unknown count unit {unit}") + if unit != "rows" and not ( + self.source.is_contiguous() and self.destination.is_contiguous() + ): + raise ValueError(f"{name}: flat/byte publication must be contiguous") + self.unit = unit + self.source_range = _range(self.source) + self.destination_range = _range(self.destination) + self.row_capacity = self.source.shape[0] if self.source.ndim else 1 + itemsize = self.source.element_size() + if unit == "rows": + self.capacity = self.row_capacity + self.bytes_per_count = itemsize + for width in self.source.shape[1:]: + self.bytes_per_count *= width + self._source_view, self._destination_view = self.source, self.destination + else: + self.capacity = self.source.numel() * (itemsize if unit == "bytes" else 1) + self.bytes_per_count = 1 if unit == "bytes" else itemsize + src, dst = self.source.reshape(-1), self.destination.reshape(-1) + if unit == "bytes": + src, dst = src.view(torch.uint8), dst.view(torch.uint8) + self._source_view, self._destination_view = src, dst + self._epoch = -1 + self._writable = False + self._single = None + # At most four private prefix pairs per binding, with FIFO eviction. + self._prefixes = {} + + def acquire_write(self, *, republish_reason=None): + """Wait before rewriting a borrowed source, never afterwards.""" + self.owner._check_active() + if self._epoch == self.owner.epoch and not _reason(republish_reason): + raise PublicationError( + f"{self.name}: repeated write needs republish_reason" + ) + if not self._writable: + self.owner._wait_sources() + self._writable = True + return self.source + + def _validate(self, count, reason): + if ( + self.buffer.cpu is not self.source + or self.buffer.gpu is not self.destination + ): + raise PublicationError( + f"{self.name}: binding changed; rebuild the buffer and owner" + ) + if self._epoch == self.owner.epoch and not _reason(reason): + raise PublicationError( + f"{self.name}: repeated publish needs republish_reason" + ) + if not self._writable: + raise PublicationError( + f"{self.name}: acquire_write before modifying source" + ) + return _count(count, self.capacity, self.name) + + def _copy(self, count): + if count == self.capacity: + self._destination_view.copy_(self._source_view, non_blocking=True) + elif count: + prefix = self._prefixes.get(count) + if prefix is None: + prefix = (self._source_view[:count], self._destination_view[:count]) + if len(self._prefixes) == 4: + del self._prefixes[next(iter(self._prefixes))] + self._prefixes[count] = prefix + prefix[1].copy_(prefix[0], non_blocking=True) + + def copy_to_gpu(self, n=None, *, republish_reason=None): + # Legacy n always means rows, including on a flat registered region. + if n is None: + count = self.capacity + else: + rows = _count(n, self.row_capacity, self.name) + count = ( + rows + if self.unit == "rows" + else ( + rows * self.capacity // self.row_capacity + if self.row_capacity + else 0 + ) + ) + self._single.publish((count,), republish_reasons=(republish_reason,)) + return self.destination if n is None else self.destination[:n] + + +class PublicationGroup: + """Fixed members validated together before any direct or kernel enqueue.""" + + def __init__(self, owner, name, members): + self.owner, self.name, self.members = owner, name, members + self._single_member = members[0] if len(members) == 1 else None + # CPU-only producer scratch; it is never borrowed by the GPU. + self.counts = [None] * len(members) + self.indices = {member.name: i for i, member in enumerate(members)} + self._buffer_indices = {member.buffer: i for i, member in enumerate(members)} + self._counts = [None] * len(members) + self._header_busy = False + self.transport = "direct" + self.fallback_reason = None + self._backend = None + + def set_count(self, buffer, count): + """Record a producer's count without copying data or granting write access. + + This only fills CPU count scratch. The producer must acquire sources + before writing them; publish validates the entire group before enqueue. + Lookup by buffer identity keeps binding names private to registration. + """ + self.counts[self._buffer_indices[buffer]] = count + + def check_writable(self): + """Check fresh sources before a producer writes any group member. + + The owner's begin() already acquired these regions. Check its common + thread/device/stream once, then each binding without another GPU API + query. This does not grant republish permission or wait for a reused + source; explicit rewrites still use BufferPublication.acquire_write. + """ + self.owner._check_active() + for binding in self.members: + binding._validate(0, None) + + def use_transport(self, transport): + """Choose direct or packed at initialization, before sources are used.""" + if self.owner._state != "idle": + raise PublicationError("choose the transport during initialization") + if transport == "packed": + if self.owner.device.type != "cuda": + backend, reason = None, "packing requires GPU destinations" + else: + from atom.utils.packed_h2d import PackedCopy + + backend, reason = PackedCopy.create(self.members, self.owner.device) + elif transport == "direct": + backend, reason = None, None + else: + raise ValueError("H2D transport must be direct or packed") + self._backend, self.fallback_reason = backend, reason + self.transport = "packed" if backend is not None else "direct" + return self.transport + + def publish(self, counts, *, republish_reasons=None): + self.owner._check_active() + backend = self._backend + if len(counts) != len(self.members): + raise ValueError(f"{self.name}: one count required per member") + if republish_reasons is not None and len(republish_reasons) != len(counts): + raise ValueError(f"{self.name}: one reason entry required per member") + # Legacy wrapper uploads use a fixed single-member group. Preserve its + # checks, ledger and failure semantics without three general loops. + if backend is None and self._single_member is not None: + count = counts[0] + if count is None: + self._counts[0] = None + return + binding = self._single_member + count = binding._validate( + count, None if republish_reasons is None else republish_reasons[0] + ) + self._counts[0] = count + binding._epoch = self.owner.epoch + binding._writable = count == 0 + try: + if count: + binding._copy(count) + except BaseException: + self.owner.fail() + raise + return + active_members = 0 + for i, (binding, count) in enumerate(zip(self.members, counts)): + self._counts[i] = ( + None + if count is None + else binding._validate( + count, None if republish_reasons is None else republish_reasons[i] + ) + ) + active_members += self._counts[i] is not None and self._counts[i] > 0 + if backend is not None and active_members > 1 and self._header_busy: + raise PublicationError( + "transport counts are in flight; acquire sources first" + ) + # No GPU-visible descriptor or ledger mutation until the WHOLE group + # validates. Explicit zero counts count; omitted members do not. + for binding, count in zip(self.members, self._counts): + if count is not None: + binding._epoch = self.owner.epoch + binding._writable = count == 0 + try: + if backend is not None and active_members: + # Direct fallback leaves any earlier packed read in flight. + self._header_busy |= backend.submit(self._counts) + elif backend is None: + for binding, count in zip(self.members, self._counts): + if count: + binding._copy(count) + except BaseException: + self.owner.fail() + raise diff --git a/atom/utils/packed_h2d.py b/atom/utils/packed_h2d.py new file mode 100644 index 0000000000..3f1252e466 --- /dev/null +++ b/atom/utils/packed_h2d.py @@ -0,0 +1,92 @@ +# SPDX-License-Identifier: MIT +"""Pack metadata into one pinned arena, upload once, scatter to stable addresses.""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def scatter_packed_bytes(arena, destinations, BLOCK: tl.constexpr): + member = tl.program_id(0) + header = arena.to(tl.pointer_type(tl.int64)) + start = tl.load(header + 2 * member) + length = tl.load(header + 2 * member + 1) + destination = tl.load(destinations + member).to(tl.pointer_type(tl.uint8)) + offset = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) + values = tl.load(arena + start + offset, offset < length, other=0) + tl.store(destination + offset, values, offset < length) + + +class PackedCopy: + """Persistent byte storage; the publication owner fences reuse of the arena.""" + + @classmethod + def create(cls, members, device): + if device.type != "cuda" or len(members) < 2: + return None, "packing requires multiple GPU destinations" + if any( + not b.source.is_contiguous() or not b.destination.is_contiguous() + for b in members + ): + return None, "strided regions use checked direct" + result = cls() + result.members = members + result.header_bytes = 16 * len(members) # int64 (offset, byte count) + capacity = result.header_bytes + sum( + b.capacity * b.bytes_per_count for b in members + ) + result.host = torch.empty( + capacity, dtype=torch.uint8, device="cpu", pin_memory=True + ) + result.device = torch.empty(capacity, dtype=torch.uint8, device=device) + result.header = result.host[: result.header_bytes].view(torch.int64).numpy() + result.payload = memoryview(result.host.numpy()) + # Byte views avoid typed copies canonicalizing bool or casting BF16. + result.sources = tuple( + memoryview(b.source.reshape(-1).view(torch.uint8).numpy()) for b in members + ) + result.destinations = torch.tensor( + [b.destination.data_ptr() for b in members], + dtype=torch.int64, + device=device, + ) + result.prefixes = {} + result.launch = scatter_packed_bytes + result.kernel = None + return result, None + + def submit(self, counts): + """Enqueue copies and return whether the packed arena is borrowed.""" + active = [i for i, count in enumerate(counts) if count] + if len(active) == 1: + # A single live member already takes one DMA; no packing/scatter. + i = active[0] + self.members[i]._copy(counts[i]) + return False + offset = self.header_bytes + largest = 0 + for i, (binding, count) in enumerate(zip(self.members, counts)): + length = 0 if count is None else count * binding.bytes_per_count + self.header[2 * i] = offset + self.header[2 * i + 1] = length + if length: + self.payload[offset : offset + length] = self.sources[i][:length] + offset += length + largest = max(largest, length) + prefix = self.prefixes.get(offset) + if prefix is None: + prefix = self.host[:offset], self.device[:offset] + if len(self.prefixes) == 4: + del self.prefixes[next(iter(self.prefixes))] + self.prefixes[offset] = prefix + # Header and payload share this single asynchronous H2D. + prefix[1].copy_(prefix[0], non_blocking=True) + grid = (len(self.members), (largest + 1023) // 1024, 1) + if self.kernel is None: + self.kernel = self.launch[grid](self.device, self.destinations, BLOCK=1024) + else: + # Both argument storages and their layouts are fixed for this + # binding. Reuse its compiled kernel; counts and grid stay dynamic. + self.kernel[grid](self.device, self.destinations) + return True diff --git a/docs/environment_variables.md b/docs/environment_variables.md index a19dc08f0b..62963b04da 100644 --- a/docs/environment_variables.md +++ b/docs/environment_variables.md @@ -2,6 +2,12 @@ This document describes the environment variables used in the ATOM project. +## Metadata H2D + +| Variable | Type | Default | Description | +|----------|------|---------|-------------| +| **ATOM_H2D_BACKEND** | str | `direct` | `packed` combines forward metadata into one H2D and GPU scatter per consumer group. `direct` copies each member separately. Both preserve source reuse gates and full cudagraph padding. Set before starting the runner. See [metadata publication](h2d_publication.md). | + ## Data parallelism | Variable | Type | Default | Description | diff --git a/docs/h2d_publication.md b/docs/h2d_publication.md new file mode 100644 index 0000000000..f12c0f94bc --- /dev/null +++ b/docs/h2d_publication.md @@ -0,0 +1,134 @@ +# Packed metadata publication + +Set `ATOM_H2D_BACKEND=packed` before starting the runner to combine small +metadata uploads. The default, `direct`, copies each member separately. +Packing is selected once for eligible groups; noncontiguous bindings use direct +copies. A single active member also uses a direct copy to avoid an extra kernel. +It neither borrows the packed arena nor releases an earlier packed read. +A producer group fully covered by an eligible larger group keeps checked +direct copies for standalone callers, without a duplicate packed arena. +Producers fill the final group's counts directly; compression plans use +`set_count(buffer, count)` without accessing buffer binding internals. + +## Consumer boundaries + +| Group | Data | First GPU consumer | +| --- | --- | --- | +| `token_inputs` | Sampling parameters, padded query prefix, input IDs, deferred source indices and speculative verification indices | Token assembly, then draft ID gather | +| `prefill_inputs` | Prefill attention mirrors and ordinary positions | Attention preparation | +| `mha_decode` | Slots, context lengths, block tables, KV prefix and positions | KV-index generation | +| `v41_metadata` | Compression/write plans, state slots, positions, token batch IDs, visibility and changed block tables | Step indptr generation | + +V4.1 ordinary DSpark decode uses two H2Ds and two scatters: token inputs and +V4.1 metadata. All graph padding is included. Other attention builders, MRoPE +and TBO can require additional groups at their own consumer boundaries. +The shared table state records publication only after the upload succeeds. + +Engram tentative preparation writes cursor candidates together with its +snapshot, preserving every accepted prefix. Committed in-place cursor updates +remain ordered after the snapshot. The sampler corrects greedy rows without +reading a GPU mask back to the CPU. Uniform top-k/top-p filters stay as CPU +scalars through ordinary and speculative sampling; only per-request filters +are uploaded. This preserves AITER scalar dispatch without a GPU readback. +The runner restores an implicit CPU default after GPU initialization to avoid +per-call DeviceContext dispatch. MRoPE decode planes use the running token +width as their axis stride before their first publication, for every backend. +V4 TBO consumers retain views of their own ubatch buffers; CPU query lengths +also supply the maximum query width without a device reduction or readback. + +## Lifetime and publication contract + +Each existing runner slot owns its pinned sources, packed arenas and completion +event. `begin()` waits for source reuse before any producer writes and starts +one forward epoch. `finish()` records completion after preparation; it does not +wait for GPU execution. PP rotates the existing slots. Late TBO preparation +resumes the same epoch and records completion again after its uploads. + +Buffers are declared with a group and count unit (`rows`, `elements` or `bytes`) +at allocation. Initialization binds their fixed source/destination storage and +checks overlaps. Producer counts include semantic padding. `None` omits a +member; zero is an explicit empty publication. A group validates all counts, +owner thread/device/stream and duplicate publications before submitting work. +Sampling, input IDs, query prefixes, speculative indices and common attention +producers declare their source groups with `@h2d_producer("group", ...)`. +The shared decorator checks the current slot before entering the producer, including when its upload uses a +combined group. Rejected producer reentry leaves earlier sources intact. +A deliberate second publication supplies a reason and acquires its source +before rewriting it. Enqueue failures invalidate the owner and retain storage. + +The packed arena contains int64 offset/count pairs and byte payloads in one DMA. +The scatter preserves bit patterns and writes only the published ranges into +the original GPU addresses. It reuses its compiled kernel and fixed destination +table. No new synchronization event or GPU scalar read is added by packing. + +Publication runs before actual graph capture/replay on the consumer stream. +Bound metadata storage stays fixed; rebuilding those buffers requires the +caller's existing consumer drain and graph rebuild. This API does not manage +checkpoint pools, bulk embeddings, model weights or KV offload storage. + +## Verification + +Run GPU publication and producer tests with `RUN_H2D_GPU_TESTS=1`. +The suites are organized by contract and consumer, with reentry regressions +kept beside the producers they protect: + +| Test file | Coverage | +| --- | --- | +| `tests/test_h2d_publication.py` | Ownership, atomic validation, stream/device identity and graph capture | +| `tests/test_packed_h2d.py` | Byte-preserving DMA/scatter, arena reuse and direct fallback | +| `tests/test_h2d_runner_publication.py` | Sampling, token inputs, query prefixes and speculative indices | +| `tests/test_h2d_attention_publication.py` | Common attention consumers, page maps and PP/TBO buffers | +| `tests/test_h2d_v4_publication.py` | V4/V4.1 plans, state, step metadata and dummy isolation | +| `tests/test_h2d_v4_indexer_publication.py` | V4 indexer, PCP, TBO and Opus | +| `tests/test_h2d_draft_publication.py` | Draft and GDN consumers | + +These suites check delayed source reuse, exact values, untouched tails, changing +batch sizes, graph padding and first consumers. CPU page-table state transitions +remain in `tests/test_shared_block_tables.py`, runnable without GPU opt-in. + +Use the server's torch profiler endpoints for model traces. On ROCm, correlate +copies with HIP runtime `kind=1`; displayed memcpy names can mislabel H2D as DtoD. +Measure from preparation GPU work through the first model kernel as well as +counting copies. Fewer copies alone do not establish a universal TTFT benefit. +Experiment logs and one-off profiling tools are maintained outside this PR. + +## Shared block-table preparation + +`block_table_state(buffer)` attaches page-map state to the existing physical +`CpuGpuBuffer`. Common attention, MHA/MLA, V4 and V4.1 share this preparation +and publication policy. Derived attention backends inherit it. Each PP slot +and each TBO buffer has its own state; this optimization is independent of TP, +DP or PP topology. The worker RPC decoder preserves versioned rows for the +same reason, rather than rebuilding page-map change information in attention. + +A source `BlockTable` carries an append-lineage version. Ordered version and +length vectors identify the batch mapping without reading page IDs. Appending retains the version; modifying or deleting existing IDs draws +a new version. The decoder copies a growing row while preserving that lineage, +so older batches remain unchanged. Independent copies draw new versions. +Unversioned array/list callers compare against the pinned snapshot instead. + +The pinned page table is the CPU snapshot: preparation retains no source +array views or intermediate ndarray/tuple copies. A hit reads only row versions +and lengths. An append checks and copies the added IDs; a changed mapping +updates only changed rows. Full-width rows are copied without clearing the +bytes they replace; fully replaced batches clear exposed tails in one bytewise +pass. The fixed destination view is reused, and no source view is retained. +Optional page-range validation follows each append lineage across reordering, +with the pool limit as part of the proof. Single-page suffixes use scalar bounds +checks. Request spans still check their required number of pages every step. All rows +are validated and source ownership acquired before any host write. + +CPU preparation and successful GPU publication have separate revisions. +Prefill can prepare CPU slots without uploading a table. Publication reuses the +GPU table only when its destination, prepared revision and requested row count +match. Changed tables still use one bulk publication; this does not add partial +H2D transfers. A combined group omits an unchanged table but consumers retain +its GPU view. Failed publication never marks the new revision as published. +Padding, TBO slices and capture table preparation use the same state. Direct +writes to the CPU/GPU table outside these entry points invalidate this contract. + +V4.1 `RequestSpan` describes only the request interval and state slot; +`begin_step(..., block_tables=rows)` receives the page mappings separately. +Runtime dummy PAGE/state storage remains private. Graph capture continues to +bind serving storage, as required by its fixed addresses. This change does not +alter either cache-allocation policy. diff --git a/tests/attentions/deepseek_v41/helpers.py b/tests/attentions/deepseek_v41/helpers.py index 3d6cd7c54d..f968fea097 100644 --- a/tests/attentions/deepseek_v41/helpers.py +++ b/tests/attentions/deepseek_v41/helpers.py @@ -1,4 +1,7 @@ # SPDX-License-Identifier: MIT +from dataclasses import dataclass + +from atom.model_ops.attentions.deepseek_v41.metadata import RequestSpan from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry @@ -83,3 +86,48 @@ def metadata_buffers(batch_size, tokens, blocks, device="cpu", geometry=None): ) ) return buffers + + +# Test cases bundle a request's span and its page mapping for readability. +# Production metadata receives them separately and retains only RequestSpan. + + +@dataclass(frozen=True) +class PagedRequest(RequestSpan): + block_ids: tuple[int, ...] + + @property + def span(self): + return RequestSpan( + self.request_id, self.position, self.offset, self.length, self.slot + ) + + +def begin_step(cache, requests, **kwargs): + requests = tuple(requests) + return cache.begin_step( + [request.span for request in requests], + block_tables=[request.block_ids for request in requests], + **kwargs, + ) + + +def prepare_step(requests, device, **kwargs): + from atom.model_ops.attentions.deepseek_v41.metadata import prepare_batch_step + + return prepare_batch_step( + [request.span for request in requests], + device, + block_tables=[request.block_ids for request in requests], + **kwargs, + ) + + +def publish_tables(buffer, requests, running_bs): + from atom.utils.block_tables import block_table_state + + return ( + block_table_state(buffer) + .prepare([request.block_ids for request in requests], pad_to=running_bs) + .publish(running_bs) + ) diff --git a/tests/attentions/deepseek_v41/test_cache.py b/tests/attentions/deepseek_v41/test_cache.py index 8da48f0a29..40f0776087 100644 --- a/tests/attentions/deepseek_v41/test_cache.py +++ b/tests/attentions/deepseek_v41/test_cache.py @@ -7,6 +7,8 @@ import pytest import torch +from tests.attentions.deepseek_v41.helpers import PagedRequest, begin_step + pytest.importorskip("aiter", reason="the paged cache and V4 kernels reach AITER") from atom.model_engine.page_unit_checkpoint import ( @@ -18,7 +20,6 @@ PagedAttentionCache, ) from atom.model_ops.attentions.deepseek_v41.checkpoints import StateCopies -from atom.model_ops.attentions.deepseek_v41.metadata import RequestSpan from atom.models.deepseek_v41.config import AttentionMode, LayerAttentionSpec from tests.attentions.deepseek_v41.helpers import geometry @@ -71,11 +72,11 @@ def test_checkpoint_fork_rollback_relocation_and_slot_reuse( copies.relocate([(1, 0), (0, 1)]) torch.testing.assert_close(copies.entry(0), original, rtol=0, atol=0) torch.testing.assert_close(copies.entry(1), old_zero, rtol=0, atol=0) - span = RequestSpan(27, 3, 0, 1, 0, (32, 33)) - step = cache.begin_step([span]) + span = PagedRequest(27, 3, 0, 1, 0, (32, 33)) + step = begin_step(cache, [span]) np.testing.assert_array_equal(cache.prepare_state(step), [[19, -1, 27]]) # Recycled slot begins at zero and drops every old state field. - fresh = cache.begin_step([replace(span, request_id=28, position=0)]) + fresh = begin_step(cache, [replace(span, request_id=28, position=0)]) np.testing.assert_array_equal(cache.prepare_state(fresh), [[-1, -1, -1]]) assert cache.state.view("window")[:, 0].count_nonzero() == 0 assert cache.state.view("compress_kv")[:, 0].count_nonzero() == 0 @@ -121,10 +122,10 @@ def test_a_rows_prefix_slice_is_exactly_as_long_as_what_gets_written(ratio, topk # where the writer's per-row count and the scan's can disagree, and the # long one is what keeps `topk` from being the only bound in play. spans = ( - RequestSpan(1, 3, 0, 1, 0, (0, 1, 2)), - RequestSpan(2, 200, 1, 1, 1, tuple(range(3, 16))), + PagedRequest(1, 3, 0, 1, 0, (0, 1, 2)), + PagedRequest(2, 200, 1, 1, 1, tuple(range(3, 16))), ) - step = cache.begin_step(spans, running_bs=2, running_tokens=2, max_q_len=1) + step = begin_step(cache, spans, running_bs=2, running_tokens=2, max_q_len=1) assert step.decode visible = (step.positions + 1) // ratio # The selection the scorers emit: ascending ids, `-1` past `min(visible, @@ -164,10 +165,10 @@ def test_one_grouped_build_writes_what_the_per_layer_builds_would(ratio): ) cache = PagedAttentionCache(geo, 32, 4, "cuda") spans = ( - RequestSpan(1, 3, 0, 1, 0, (0, 1, 2)), - RequestSpan(2, 200, 1, 1, 1, tuple(range(3, 16))), + PagedRequest(1, 3, 0, 1, 0, (0, 1, 2)), + PagedRequest(2, 200, 1, 1, 1, tuple(range(3, 16))), ) - step = cache.begin_step(spans, running_bs=2, running_tokens=2, max_q_len=1) + step = begin_step(cache, spans, running_bs=2, running_tokens=2, max_q_len=1) assert step.decode selection = torch.where( torch.arange(8, device="cuda") < ((step.positions + 1) // ratio)[:, None], @@ -217,7 +218,7 @@ def test_graph_plan_sentinel_rows_do_not_write_to_live_pages(): geo = V41PoolGeometry(2, ((0, 2), (1, 2)), 32, 4, 512, 32, speculative_tokens=1) cache = PagedAttentionCache(geo, 6, 2, "cpu") running_bs, max_q_len = 2, 2 - spans = (RequestSpan(1, 4, 0, 2, 0, (3, 5)),) + spans = (PagedRequest(1, 4, 0, 2, 0, (3, 5)),) plans = make_compress_plans( np.asarray([2], dtype=np.int32), np.asarray([6], dtype=np.int32), @@ -232,8 +233,12 @@ def test_graph_plan_sentinel_rows_do_not_write_to_live_pages(): max_q_len=max_q_len, extra_write=1, ) - step = cache.begin_step( - spans, running_bs=running_bs, running_tokens=running_bs * max_q_len, plans=plans + step = begin_step( + cache, + spans, + running_bs=running_bs, + running_tokens=running_bs * max_q_len, + plans=plans, ) plan = step.plans[2] # The capacity, which is what the kernel's grid and every row derived from @@ -332,17 +337,17 @@ def test_a_deferred_prepare_state_names_the_stale_slot_one_step_later( cache = PagedAttentionCache(geo, 8, 4, device) cache.cursor[0] = torch.tensor([3, 19, -1, 27], device=device) - good = RequestSpan(27, 3, 0, 1, 0, (0, 1)) - step = cache.begin_step([good]) + good = PagedRequest(27, 3, 0, 1, 0, (0, 1)) + step = begin_step(cache, [good]) assert cache.prepare_state(step, histories=False) is None # Same slot, but the scheduler believes it is four tokens further on than # the cursor says. Deferred, so this call is the one that ships the rows. - stale = RequestSpan(28, 7, 0, 1, 0, (0, 1)) - later = cache.begin_step([stale]) + stale = PagedRequest(28, 7, 0, 1, 0, (0, 1)) + later = begin_step(cache, [stale]) assert cache.prepare_state(later, histories=False) is None with pytest.raises(ValueError, match="Request 28 needs state at 7"): - cache.prepare_state(cache.begin_step([good]), histories=False) + cache.prepare_state(begin_step(cache, [good]), histories=False) @pytest.mark.parametrize("device", ["cpu", "cuda"]) @@ -353,11 +358,11 @@ def test_a_deferred_probe_skips_the_slot_its_own_step_resets(small_config, devic geo = replace(geometry(small_config), window_size=32) cache = PagedAttentionCache(geo, 8, 4, device) cache.cursor[0] = torch.tensor([91, 19, -1, 27], device=device) - fresh = RequestSpan(31, 0, 0, 1, 0, (0, 1)) - cache.prepare_state(cache.begin_step([fresh]), histories=False) + fresh = PagedRequest(31, 0, 0, 1, 0, (0, 1)) + cache.prepare_state(begin_step(cache, [fresh]), histories=False) assert cache.cursor[0, 0] == 0 # The probe carries slot 0's pre-reset row; position 0 is what excludes it. - cache.prepare_state(cache.begin_step([replace(fresh, request_id=32)])) + cache.prepare_state(begin_step(cache, [replace(fresh, request_id=32)])) @pytest.mark.skipif(not torch.cuda.is_available(), reason="ROCm GPU required") diff --git a/tests/attentions/deepseek_v41/test_metadata.py b/tests/attentions/deepseek_v41/test_metadata.py index f6f9573574..b009d96265 100644 --- a/tests/attentions/deepseek_v41/test_metadata.py +++ b/tests/attentions/deepseek_v41/test_metadata.py @@ -6,14 +6,17 @@ import pytest import torch +from tests.attentions.deepseek_v41.helpers import ( + PagedRequest, + begin_step, + publish_tables, +) + pytest.importorskip("aiter", reason="the V4.1 backend and cache reach AITER") from atom.model_ops.attentions.deepseek_v41.backend import DeepseekV41MetadataBuilder from atom.model_ops.attentions.deepseek_v41.cache import PagedAttentionCache -from atom.model_ops.attentions.deepseek_v41.metadata import ( - RequestSpan, - visible_buffer_name, -) +from atom.model_ops.attentions.deepseek_v41.metadata import visible_buffer_name from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry from atom.utils import CpuGpuBuffer from tests.attentions.deepseek_v41.helpers import metadata_buffers @@ -27,7 +30,8 @@ def staged_step(cache, requests, *, buffers, running_bs, running_tokens): len(requests), running_bs, ) - return cache.begin_step( + return begin_step( + cache, requests, buffers=buffers, running_bs=running_bs, @@ -44,12 +48,12 @@ def test_staged_metadata_refreshes_reordered_ragged_and_empty_batches(device): cache = PagedAttentionCache(geo, 8, 4, device) buffers = metadata_buffers(4, 8, 4, device, geo) pointers = {name: value.gpu.data_ptr() for name, value in buffers.items()} - first = RequestSpan(17, 1, 0, 3, 3, (5, 1)) - second = RequestSpan(24, 4, 3, 1, 1, (2, 7)) + first = PagedRequest(17, 1, 0, 3, 3, (5, 1)) + second = PagedRequest(24, 4, 3, 1, 1, (2, 7)) for requests in ( (first, second), - (RequestSpan(24, 5, 0, 2, 1, (2, 7)), RequestSpan(17, 4, 2, 1, 3, (5, 1))), - (RequestSpan(9, 0, 0, 1, 0, (6,)),), + (PagedRequest(24, 5, 0, 2, 1, (2, 7)), PagedRequest(17, 4, 2, 1, 3, (5, 1))), + (PagedRequest(9, 0, 0, 1, 0, (6,)),), (), ): step = staged_step( @@ -111,7 +115,7 @@ def test_engram_rows_are_staged_for_the_width_not_for_the_tokens(engram): buffers = metadata_buffers(4, 8, 4, "cpu", geo) step = staged_step( cache, - (RequestSpan(17, 1, 0, 3, 3, (5, 1)),), + (PagedRequest(17, 1, 0, 3, 3, (5, 1)),), buffers=buffers, running_bs=2, running_tokens=6, @@ -154,6 +158,8 @@ def test_engram_rows_are_staged_for_the_width_not_for_the_tokens(engram): # A synthetic batch hashes on the host: the cursor the device path # would read belongs to whoever owns these slots, not to this one. "batch": None, + "cursor_positions": None, + "cursor_out": None, } @@ -162,7 +168,7 @@ def test_v4_window_write_graph_reads_updated_requests_without_recapture(): geo = V41PoolGeometry(1, ((0, 2),), 32, 4, 512, 32) cache = PagedAttentionCache(geo, 8, 4, "cuda") buffers = metadata_buffers(2, 2, 2, "cuda", geo) - first = (RequestSpan(17, 0, 0, 1, 1, (0, 1)), RequestSpan(24, 2, 1, 1, 3, (2, 3))) + first = (PagedRequest(17, 0, 0, 1, 1, (0, 1)), PagedRequest(24, 2, 1, 1, 3, (2, 3))) step = staged_step(cache, first, buffers=buffers, running_bs=2, running_tokens=2) values = torch.randn(1, 2, 512, device="cuda", dtype=torch.bfloat16) stream = torch.cuda.Stream() @@ -175,8 +181,8 @@ def test_v4_window_write_graph_reads_updated_requests_without_recapture(): torch.cuda.current_stream().wait_stream(stream) window = cache.state.view("window")[0] for spans in ( - (RequestSpan(24, 3, 0, 1, 3, (2, 3)), RequestSpan(17, 1, 1, 1, 1, (0, 1))), - (RequestSpan(17, 4, 0, 1, 0, (4, 5)), RequestSpan(24, 6, 1, 1, 2, (6, 7))), + (PagedRequest(24, 3, 0, 1, 3, (2, 3)), PagedRequest(17, 1, 1, 1, 1, (0, 1))), + (PagedRequest(17, 4, 0, 1, 0, (4, 5)), PagedRequest(24, 6, 1, 1, 2, (6, 7))), ): staged_step(cache, spans, buffers=buffers, running_bs=2, running_tokens=2) window.zero_() @@ -242,8 +248,8 @@ def test_visible_rows_are_the_bound_every_indexer_layer_reads(ratio, positions_d buffers = metadata_buffers(4, 8, 4, "cpu", geo) buffers["positions"] = CpuGpuBuffer(8, dtype=positions_dtype, device="cpu") requests = ( - RequestSpan(17, 1, 0, 3, 3, (5, 1)), - RequestSpan(24, 9, 3, 1, 1, (2, 7)), + PagedRequest(17, 1, 0, 3, 3, (5, 1)), + PagedRequest(24, 9, 3, 1, 1, (2, 7)), ) step = staged_step(cache, requests, buffers=buffers, running_bs=4, running_tokens=8) # Armed: ragged, and padded past the batch, so the two layouts really do @@ -267,8 +273,6 @@ def test_block_table_upload_only_when_mapping_changes(device, monkeypatch): pytest.skip("ROCm GPU required") from dataclasses import replace - from atom.model_ops.attentions.deepseek_v41.metadata import _publish_block_tables - tables = CpuGpuBuffer( 4, 8, dtype=torch.int32, device=device, pin_memory=device != "cpu" ) @@ -280,8 +284,8 @@ def counted_copy(n=None): return original_copy(n) monkeypatch.setattr(tables, "copy_to_gpu", counted_copy) - a = RequestSpan(1, 12, 0, 1, 0, (3, 5)) - b = RequestSpan(2, 22, 1, 1, 1, (2, 6)) + a = PagedRequest(1, 12, 0, 1, 0, (3, 5)) + b = PagedRequest(2, 22, 1, 1, 1, (2, 6)) cases = [ ((a, b), 4, True), ((replace(a, position=13), replace(b, position=23)), 4, False), @@ -295,7 +299,7 @@ def counted_copy(n=None): ] for requests, running_bs, changed in cases: before = len(uploads) - result = _publish_block_tables(tables, requests, running_bs) + result = publish_tables(tables, requests, running_bs) assert len(uploads) - before == int(changed) expected = [list(s.block_ids) + [0] * (8 - len(s.block_ids)) for s in requests] expected += [[0] * 8 for _ in range(running_bs - len(requests))] @@ -303,6 +307,6 @@ def counted_copy(n=None): assert result.data_ptr() == tables.gpu.data_ptr() # Replacing the device allocation invalidates reuse even for identical rows. tables.gpu = torch.full_like(tables.gpu, -1) - _publish_block_tables(tables, (a, b), 4) + publish_tables(tables, (a, b), 4) assert len(uploads) == sum(c[2] for c in cases) + 1 assert tables.gpu[0, :2].tolist() == [3, 5] diff --git a/tests/attentions/deepseek_v41/test_runtime_contract.py b/tests/attentions/deepseek_v41/test_runtime_contract.py index 3c62fe6fcb..7092f1dbbf 100644 --- a/tests/attentions/deepseek_v41/test_runtime_contract.py +++ b/tests/attentions/deepseek_v41/test_runtime_contract.py @@ -9,9 +9,9 @@ import pytest import torch -from atom.model_ops.attentions.deepseek_v41.metadata import RequestSpan from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry from atom.models.deepseek_v41.config import normalize_hf_config, validate_runtime_config +from tests.attentions.deepseek_v41.helpers import PagedRequest, begin_step def _runtime_pieces(): @@ -125,7 +125,7 @@ def test_empty_rank_padding_has_no_cache_writes(monkeypatch): cache = PagedAttentionCache(geo, 4, 2, "cpu") cache.backing.fill_(57) before = cache.backing.clone() - step = cache.begin_step([]) + step = begin_step(cache, []) metadata = SimpleNamespace( step=step, cache=cache, @@ -163,7 +163,7 @@ def test_a_forward_reads_nothing_the_forward_before_it_selected(monkeypatch): geo = V41PoolGeometry(2, ((1, 2),), 32, 4, 512, 32) cache = PagedAttentionCache(geo, 4, 2, "cpu") - step = cache.begin_step([RequestSpan(0, 0, 0, 1, 0, (0,))], plans={}) + step = begin_step(cache, [PagedRequest(0, 0, 0, 1, 0, (0,))], plans={}) memos = {name: getattr(step, name) for name in ("selected", "candidates")} for name, memo in memos.items(): memo["what the last forward worked out"] = name diff --git a/tests/model_ops/engram/test_overlap.py b/tests/model_ops/engram/test_overlap.py index 52bd85f438..5a065c5c5a 100644 --- a/tests/model_ops/engram/test_overlap.py +++ b/tests/model_ops/engram/test_overlap.py @@ -20,6 +20,7 @@ from atom.model_ops.engram.device.hashing import ( EngramHashTables, + engram_cursor_rows_reference, engram_row_indices_reference, engram_snapshot, engram_snapshot_indices, @@ -87,6 +88,67 @@ def test_snapshot_keeps_old_history_and_padding(config, lengths): assert torch.all(out[count:] == -1) +@pytest.mark.parametrize("config", [tiny_config(), EngramConfig.from_hf(V41_FLASH)]) +@pytest.mark.parametrize("lengths", [[1, 1], [1, 6, 3], [6] * 8, [129, 3, 1]]) +@pytest.mark.parametrize("replay", [False, True]) +def test_snapshot_stages_all_cursor_prefixes_without_mutating_history( + config, lengths, replay +): + mapping = build(config) + tables = EngramHashTables.from_mapping(mapping, torch.device("cuda")) + batch, tokens, histories, masks = make_batch(mapping, lengths, 7) + count = sum(lengths) + starts = np.arange(len(lengths)) * 31 + 2 + positions = torch.tensor( + np.concatenate( + [np.arange(start, start + n) for start, n in zip(starts, lengths)] + ), + dtype=torch.int64, + device="cuda", + ) + # Unscheduled prefixes and requests must retain their sentinels. Snapshot + # padding is still written as -2, including across a CTA boundary. + snapshot = torch.empty(count + 5, tables.ngram, dtype=torch.int64, device="cuda") + candidates = torch.full( + (len(lengths) + 1, max(lengths) + 2, tables.ngram), + -99, + dtype=torch.int64, + device="cuda", + ) + original = batch.history.clone() + + def prepare(): + engram_snapshot( + tables, batch, snapshot, cursor_positions=positions, cursor_out=candidates + ) + + prepare() + if replay: + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + prepare() + # Fixed addresses, changed data: replay must not reuse a cached result. + batch.compressed.fill_(-1) + candidates.fill_(-99) + graph.replay() + masks = [np.zeros(n, dtype=bool) for n in lengths] + expected = engram_cursor_rows_reference(mapping, tokens, histories, starts, masks) + for i, (length, rows) in enumerate(zip(lengths, expected)): + np.testing.assert_array_equal(candidates[i, :length].cpu().numpy(), rows) + assert torch.all(candidates[i, length:] == -99) + assert torch.all(candidates[len(lengths) :] == -99) + assert torch.equal(batch.history, original) + assert torch.all(snapshot[count:] == -2) + for layer in config.layer_ids: + out = torch.empty(count + 5, tables.heads, dtype=torch.int64, device="cuda") + engram_snapshot_indices(tables, layer, snapshot, out) + np.testing.assert_array_equal( + out[:count].cpu().numpy(), + engram_row_indices_reference(mapping, layer, tokens, histories, masks), + ) + assert torch.all(out[count:] == -1) + + def make_staging(mapping, group=None): config = mapping.config device = torch.device("cuda", torch.cuda.current_device()) diff --git a/tests/models/deepseek_v41/test_dspark_integration.py b/tests/models/deepseek_v41/test_dspark_integration.py index 3509ce38a1..e5f8251cd2 100644 --- a/tests/models/deepseek_v41/test_dspark_integration.py +++ b/tests/models/deepseek_v41/test_dspark_integration.py @@ -7,6 +7,8 @@ import pytest import torch +from tests.attentions.deepseek_v41.helpers import PagedRequest, begin_step + pytest.importorskip("aiter", reason="the draft stack builds AITER-backed layers") from torch import nn @@ -94,14 +96,17 @@ def test_draft_context_write_spans_the_forwards_width_not_its_tokens(monkeypatch while the read side gathers by absolute position regardless. """ from atom.model_ops.attentions.deepseek_v41.cache import PagedAttentionCache - from atom.model_ops.attentions.deepseek_v41.metadata import RequestSpan from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry cache = PagedAttentionCache( V41PoolGeometry(1, ((0, 2),), 32, 4, 512, 32), 8, 4, "cpu" ) - step = cache.begin_step( - [RequestSpan(0, 0, 0, 1, 0, (0,))], running_bs=2, running_tokens=2, plans={} + step = begin_step( + cache, + [PagedRequest(0, 0, 0, 1, 0, (0,))], + running_bs=2, + running_tokens=2, + plans={}, ) # Three distinct numbers, so a slice by the wrong one cannot pass. assert step.scheduled == 1 and step.width == 2 @@ -229,6 +234,8 @@ def test_decode_positions_use_accepted_prefix_and_full_reservation(): num_scheduled_tokens=(1, 3), block_tables=(tuple(range(10)), tuple(range(10, 20))), ) + # Serving publishes this before input assembly and attention preparation. + builder.publish_cu_seqlens_q(batch, SimpleNamespace(running_bs=2)) metadata, actual = builder.prepare_decode(batch, 2, 4, 3) assert [span.position for span in metadata.step.requests] == [129, 140] assert actual.tolist() == [129, 140, 141, 142] diff --git a/tests/models/deepseek_v41/test_sparse_attention.py b/tests/models/deepseek_v41/test_sparse_attention.py index 8d35c4e4cc..bd7143feef 100644 --- a/tests/models/deepseek_v41/test_sparse_attention.py +++ b/tests/models/deepseek_v41/test_sparse_attention.py @@ -4,6 +4,8 @@ import pytest import torch +from tests.attentions.deepseek_v41.helpers import PagedRequest, begin_step + pytest.importorskip("aiter", reason="the V4 kernels import the AITER runtime") from atom.model_ops.v4_kernels import ( @@ -26,7 +28,6 @@ def test_v4_bf16_counts_sink_once_for_swa_and_global(small_config, length): row double-counted would move it. """ from atom.model_ops.attentions.deepseek_v41.cache import PagedAttentionCache - from atom.model_ops.attentions.deepseek_v41.metadata import RequestSpan from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry config = small_config @@ -44,10 +45,10 @@ def test_v4_bf16_counts_sink_once_for_swa_and_global(small_config, length): cache = PagedAttentionCache(geo, 8, 2, "cuda") cache.pages.view("main_0").fill_(6) spans = ( - RequestSpan(1, 0, 0, length, 0, (0, 1)), - RequestSpan(2, 0, length, length, 1, (2, 3)), + PagedRequest(1, 0, 0, length, 0, (0, 1)), + PagedRequest(2, 0, length, length, 1, (2, 3)), ) - step = cache.begin_step(spans) + step = begin_step(cache, spans) spec = LayerAttentionSpec(0, 1, AttentionMode.FULL, 0, 0) # Index row 0 for every query row: the same row its window already holds. step.selected[0] = torch.zeros(1, step.width, 1, device="cuda", dtype=torch.int32) diff --git a/tests/models/deepseek_v41/test_speculative_state.py b/tests/models/deepseek_v41/test_speculative_state.py index ab243fbe19..cbdc1495a3 100644 --- a/tests/models/deepseek_v41/test_speculative_state.py +++ b/tests/models/deepseek_v41/test_speculative_state.py @@ -4,11 +4,12 @@ import pytest import torch +from tests.attentions.deepseek_v41.helpers import PagedRequest, begin_step + pytest.importorskip("aiter", reason="the paged cache reaches AITER") from atom.model_ops.attentions.deepseek_v41.cache import PagedAttentionCache from atom.model_ops.attentions.deepseek_v41.checkpoints import StateCopies -from atom.model_ops.attentions.deepseek_v41.metadata import RequestSpan from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry from atom.models.deepseek_v41.config import AttentionMode, LayerAttentionSpec @@ -64,8 +65,8 @@ def test_every_prefix_survives_ring_wrap_and_ragged_request_order( ) cache = PagedAttentionCache(geometry, 40, 5, device) spans = ( - RequestSpan(11, position, 0, 6, 4, tuple(range(20))), - RequestSpan(22, position - 4, 6, 3, 1, tuple(range(20, 40))), + PagedRequest(11, position, 0, 6, 4, tuple(range(20))), + PagedRequest(22, position - 4, 6, 3, 1, tuple(range(20, 40))), ) history = torch.tensor([41, -1, 43], device=device) for span in spans: @@ -80,7 +81,7 @@ def test_every_prefix_survives_ring_wrap_and_ragged_request_order( ) untouched = cache.state_bytes[0].clone() old_cursors = cache.cursor.clone() - step = cache.begin_step(spans, tentative=True) + step = begin_step(cache, spans, tentative=True) cache.prepare_state(step) expected_histories = [] accepted_lengths = (accepted, min(accepted, 3)) @@ -94,7 +95,7 @@ def test_every_prefix_survives_ring_wrap_and_ragged_request_order( ) assert torch.equal(cache.cursor, old_cursors) with pytest.raises(RuntimeError, match="Commit the accepted prefix"): - cache.begin_step(spans) + begin_step(cache, spans) copies = StateCopies.__new__(StateCopies) copies.cache = cache with pytest.raises(RuntimeError, match="Commit the accepted prefix"): @@ -113,12 +114,12 @@ def test_every_prefix_survives_ring_wrap_and_ragged_request_order( # Reorder requests after acceptance. Addressing must use each noncontiguous # STATE slot, and logical visibility must stay at 128, not the 133-row ring. next_spans = tuple( - RequestSpan( + PagedRequest( span.request_id, span.position + count, i * 2, 2, span.slot, span.block_ids ) for i, (span, count) in enumerate(reversed(list(zip(spans, accepted_lengths)))) ) - next_step = cache.begin_step(next_spans) + next_step = begin_step(cache, next_spans) cache.prepare_state(next_step) if device == "cuda": for layer in range(2): @@ -152,9 +153,9 @@ def test_tentative_state_refuses_missing_prefixes_and_a_stale_step(): geometry = V41PoolGeometry(1, ((0, 2),), 32, 128, 128, 32, speculative_tokens=5) cache = PagedAttentionCache(geometry, 1, 1, "cpu") cache.cursor[0, 0] = 3 - span = RequestSpan(1, 3, 0, 6, 0, (0,)) - stale = cache.begin_step((span,)) - step = cache.begin_step((span,), tentative=True) + span = PagedRequest(1, 3, 0, 6, 0, (0,)) + stale = begin_step(cache, (span,)) + step = begin_step(cache, (span,), tentative=True) cache.prepare_state(step) with pytest.raises(RuntimeError, match="missing"): cache.commit_tentative(step, [1]) @@ -180,8 +181,8 @@ def test_commit_moves_the_scheduled_cursors_and_no_padding_requests(): cache = PagedAttentionCache(geometry, 1, 4, "cpu") for slot in range(4): cache.cursor[slot, 0] = 3 - spans = (RequestSpan(1, 3, 0, 2, 3, (0,)), RequestSpan(2, 3, 2, 2, 1, (0,))) - step = cache.begin_step(spans, tentative=True, running_bs=3, running_tokens=6) + spans = (PagedRequest(1, 3, 0, 2, 3, (0,)), PagedRequest(2, 3, 2, 2, 1, (0,))) + step = begin_step(cache, spans, tentative=True, running_bs=3, running_tokens=6) assert step.scheduled_bs == 2 and step.slots.tolist() == [3, 1, 0] cache.prepare_state(step) for span in spans: @@ -224,8 +225,8 @@ def test_a_rejected_round_leaves_the_next_one_as_if_it_never_drafted( def round_one(length): cache = PagedAttentionCache(geometry, 20, 1, "cuda") cache.cursor[0, 0] = position - span = RequestSpan(7, position, 0, length, 0, blocks) - step = cache.begin_step((span,), tentative=True) + span = PagedRequest(7, position, 0, length, 0, blocks) + step = begin_step(cache, (span,), tentative=True) cache.prepare_state(step) cache.compress( 0, compressor, *compressor.project(hidden[:, :length]), step, rope @@ -235,8 +236,8 @@ def round_one(length): return cache def round_two(cache): - span = RequestSpan(7, position + accepted, 0, 3, 0, blocks) - step = cache.begin_step((span,)) + span = PagedRequest(7, position + accepted, 0, 3, 0, blocks) + step = begin_step(cache, (span,)) cache.prepare_state(step) latent = cache.compress( 0, compressor, *compressor.project(following), step, rope @@ -275,10 +276,10 @@ def test_block_context_read_decodes_only_the_selected_request_windows(packed): ) cache = PagedAttentionCache(geometry, 24, 4, "cuda") spans = ( - RequestSpan(1, 0, 0, 145, 3, tuple(range(10))), - RequestSpan(2, 0, 145, 142, 1, tuple(range(10, 20))), + PagedRequest(1, 0, 0, 145, 3, tuple(range(10))), + PagedRequest(2, 0, 145, 142, 1, tuple(range(10, 20))), ) - step = cache.begin_step(spans) + step = begin_step(cache, spans) cache.prepare_state(step) torch.manual_seed(95) values = torch.randn(1, 287, 512, device="cuda", dtype=torch.bfloat16) @@ -322,7 +323,7 @@ def test_verify_decode_kernel_is_causal_after_writing_the_whole_block(packed): stored = pack_rows(*quantized) if packed else keys cache.state.view("window")[0, 0, :7] = stored[:7] cache.cursor[0, 0] = 7 - step = cache.begin_step((RequestSpan(1, 7, 0, 6, 0, (0,)),), tentative=True) + step = begin_step(cache, (PagedRequest(1, 7, 0, 6, 0, (0,)),), tentative=True) cache.prepare_state(step) assert step.decode cache.write_window( diff --git a/tests/test_h2d_attention_publication.py b/tests/test_h2d_attention_publication.py new file mode 100644 index 0000000000..caee8869e9 --- /dev/null +++ b/tests/test_h2d_attention_publication.py @@ -0,0 +1,627 @@ +# SPDX-License-Identifier: MIT +"""Persistent attention producers through real H2D owners and GPU consumers.""" + +import os +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from atom.utils import CpuGpuBuffer +from atom.utils.h2d import PublicationError + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1", reason="set RUN_H2D_GPU_TESTS=1" +) + + +def make_runner( + monkeypatch, + kind="mla", + *, + pp=1, + q=1, + dcp=1, + sparse=False, + tbo=True, + base=None, + bind=True, + transport="direct", +): + from atom.model_engine.model_runner import ModelRunner + from atom.model_ops.attentions import aiter_attention, aiter_mla, backends + + monkeypatch.setenv("ATOM_H2D_BACKEND", transport) + monkeypatch.setenv("ATOM_MLA_PAGE_SIZE", "1") + monkeypatch.setenv("ATOM_USE_UNIFIED_ATTN", "0") + monkeypatch.setenv("ATOM_USE_TRITON_MLA", "0") + monkeypatch.setattr(backends, "tbo_enabled", lambda: tbo) + for module in (backends, aiter_attention, aiter_mla): + if hasattr(module, "get_tp_group"): + monkeypatch.setattr( + module, "get_tp_group", lambda: SimpleNamespace(world_size=1) + ) + if hasattr(module, "get_dcp_world_size"): + monkeypatch.setattr(module, "get_dcp_world_size", lambda: dcp) + if hasattr(module, "get_dcp_rank"): + monkeypatch.setattr(module, "get_dcp_rank", lambda: dcp - 1) + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.max_bs = 4 + runner.max_num_batched_tokens = 64 + runner.block_size = 16 + runner.kv_cache_dtype = "bf16" + runner.num_spec_tokens = q - 1 + runner.has_mla_indexer = sparse + runner.use_mrope = False + runner.enforce_eager = True + runner.arange_np = np.arange(64, dtype=np.int64) + runner.config = SimpleNamespace( + pipeline_parallel_size=pp, + max_model_len=128, + enable_tbo=tbo, + enable_tbo_decode=tbo, + kv_cache_dtype="bf16", + index_cache_dtype="fp8", + attn_prefill_chunk_size=0, + compilation_config=SimpleNamespace(static_forward_context={}), + speculative_config=SimpleNamespace(num_speculative_tokens=q - 1), + hf_config=SimpleNamespace( + num_attention_heads=16, + num_key_value_heads=1, + index_topk=8, + index_kpool=1, + indexer_compress_ratio=4, + ngram_size=2, + ple_layer_ids=[0], + eos_token_id=2, + ), + ) + runner.tokenID_processor = SimpleNamespace(num_rejected=None) + runner.forward_vars = { + name: CpuGpuBuffer(64, dtype=torch.int32, device=runner.device) + for name in ("input_ids", "decode_src") + } + runner.forward_vars["positions"] = CpuGpuBuffer( + 64, dtype=torch.int64, device=runner.device, publication_group="positions" + ) + runner.forward_vars["mtp_k"] = q - 1 + cls = base or ( + aiter_mla.AiterMLAMetadataBuilder + if kind == "mla" + else aiter_attention.AiterAttentionMetadataBuilder + ) + builder = cls(model_runner=runner) + # These tests execute the host producers and real CSR/compressed-slot + # consumers. Attention worker scheduling is covered by its own suite. + if kind == "mla": + builder.set_mla_persistent_worker_buffers = lambda *a, **kw: {} + builder._set_mla_persistent_worker_buffers_sparse_mtp = lambda *a, **kw: {} + builder._set_ubatch_mla_buffers = lambda *a, **kw: None + builder._publish_indexer_fp4_decode_schedule = lambda *a, **kw: None + builder._publish_indexer_fp4_prefill_schedule = lambda *a, **kw: None + runner.attn_metadata_builder = builder + runner._init_forward_vars_ring() + if bind: + runner._init_h2d_publication() + return runner, builder + + +def decode_batch(count, q=1, step=0): + lens = np.array([17 + 3 * i + step for i in range(count)], dtype=np.int32) + return SimpleNamespace( + total_seqs_num_decode=count, + total_tokens_num_decode=count * q, + total_seqs_num_prefill=0, + total_tokens_num_prefill=0, + total_seqs_num=count, + total_tokens_num=count * q, + context_lens=lens, + num_scheduled_tokens=np.full(count, q, dtype=np.int32), + block_tables=[ + np.asarray( + [5 + 8 * i + j for j in range((int(n) + 15) // 16)], dtype=np.int32 + ) + for i, n in enumerate(lens) + ], + last_block_num_tokens=[(int(n) - 1) % 16 + 1 for n in lens], + is_first_decode_without_local_prefill=[False] * count, + is_dummy_run=False, + ) + + +def start_decode(runner, builder, batch, q): + runner._advance_forward_vars() + runner._gate_staging_reuse() + builder.publish_cu_seqlens_q(batch, SimpleNamespace(running_bs=4)) + + +def finish(runner): + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + + +def assert_repeat_does_not_write(runner, call): + before = { + name: buf.cpu.clone() + for name, buf in runner.forward_vars.items() + if isinstance(buf, CpuGpuBuffer) + } + with pytest.raises(PublicationError, match="republish_reason"): + call() + for name, expected in before.items(): + assert torch.equal(runner.forward_vars[name].cpu, expected), name + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +@pytest.mark.parametrize("pp", [1, 2]) +@pytest.mark.parametrize( + "kind,q,dcp,sparse,tbo", + [ + ("mha", 1, 1, False, True), + ("mha", 3, 1, False, True), + ("mla", 1, 1, False, True), + ("mla", 3, 2, False, True), + ("mla", 1, 2, True, True), + ("mla", 3, 2, True, False), + ], +) +def test_decode_csr_tbo_delayed_slot_reuse( + monkeypatch, pp, kind, q, dcp, sparse, tbo, transport +): + runner, builder = make_runner( + monkeypatch, + kind, + pp=pp, + q=q, + dcp=dcp, + sparse=sparse, + tbo=tbo, + transport=transport, + ) + saved = [] + for step, count in enumerate((3, 1, 4, 2, 1, 3)): + batch = decode_batch(count, q, step) + start_decode(runner, builder, batch, q) + torch.cuda._sleep(3_000_000) + md, _ = builder.prepare_decode(batch, 4, 4 * q, q) + names = ["kv_indptr"] + if kind == "mla": + names.append("kv_last_page_lens") + if dcp > 1: + names.append("g_kv_indptr") + if sparse: + names.extend(("sparse_kv_indptr", "dcp_local_context_lens")) + for name in names: + assert ( + runner.forward_vars[name]._publication._epoch == runner.h2d_owner.epoch + ) + lengths = batch.context_lens + local = (lengths + dcp - 1 - builder.dcp_rank) // dcp + pages = (local + builder.block_size - 1) // builder.block_size + expected = [0] + np.cumsum(pages).tolist() + [int(pages.sum())] * (4 - count) + saved.append((md.kv_indptr.clone(), expected)) + # The CSR generator consumes the freshly published indptr/table on GPU. + indices = [] + for table, n in zip(batch.block_tables, pages): + indices.extend( + table[j // builder.block_ratio] * builder.block_ratio + + j % builder.block_ratio + for j in range(int(n)) + ) + saved.append((md.kv_indices[: len(indices)].clone(), indices)) + if dcp > 1: + expected_global = ( + [0] + np.cumsum(lengths).tolist() + [int(lengths.sum())] * (4 - count) + ) + saved.append((md.g_kv_indptr.clone(), expected_global)) + if sparse: + expected_local = [ + (int(n) - q + j + 1 + dcp - 1 - builder.dcp_rank) // dcp + for n in lengths + for j in range(q) + ] + [0] * ((4 - count) * q) + saved.append((md.dcp_local_context_lens.clone(), expected_local)) + effective = [min(int(n) - q + j + 1, 8) for n in lengths for j in range(q)] + expected_sparse = ( + [0] + + np.cumsum(effective).tolist() + + [sum(effective)] * ((4 - count) * q) + ) + saved.append( + ( + runner.forward_vars["sparse_kv_indptr"].gpu[: 4 * q + 1].clone(), + expected_sparse, + ) + ) + if tbo: + for ub, lo in enumerate((0, 2)): + n = max(0, min(count - lo, 2)) + p = f"ub{ub}_" + expected_ub = [v - expected[lo] for v in expected[lo : lo + 3]] + saved.append( + (runner.forward_vars[p + "kv_indptr"].gpu[:3].clone(), expected_ub) + ) + saved.append( + ( + runner.forward_vars[p + "cu_seqlens_q"].gpu[:3].clone(), + [0] + [min(j, n) * q for j in (1, 2)], + ) + ) + assert_repeat_does_not_write( + runner, + lambda count=count, lengths=lengths: builder._prepare_ubatch_decode( + count, 4, q, lengths + ), + ) + assert_repeat_does_not_write( + runner, + lambda count=count, step=step: builder.prepare_decode( + decode_batch(count, q, step + 50), 4, 4 * q, q + ), + ) + finish(runner) + torch.cuda.synchronize() + for actual, expected in saved: + assert actual.cpu().tolist() == expected + + +@pytest.mark.parametrize("kind", ["mla", "mha"]) +def test_tbo_empty_batch_and_preflight_all_sources(monkeypatch, kind): + runner, builder = make_runner(monkeypatch, kind) + runner._gate_staging_reuse() + builder._prepare_ubatch_decode(0, 4, 1, np.array([], dtype=np.int32)) + finish(runner) + torch.cuda.synchronize() + for ub in range(2): + var = runner.forward_vars + assert var[f"ub{ub}_kv_indptr"].gpu[:3].tolist() == [0, 0, 0] + assert var[f"ub{ub}_slot_mapping"].gpu[:2].tolist() == [-1, -1] + runner._gate_staging_reuse() + runner.forward_vars["ub1_kv_indptr"].copy_to_gpu(0) + assert_repeat_does_not_write( + runner, + lambda: builder._prepare_ubatch_decode(0, 4, 1, np.array([], dtype=np.int32)), + ) + finish(runner) + + +@pytest.mark.parametrize("dcp,q", [(1, 1), (2, 3)]) +def test_kimi_capture_has_one_publication_and_rejects_rewrite(monkeypatch, dcp, q): + from atom.model_ops.attentions.kimi_mla_gdn_attn import ( + KimiAiterMLAGDNMetadataBuilder, + ) + + runner, builder = make_runner(monkeypatch, q=q, dcp=dcp, tbo=False) + # Reuse production MLA allocation; Kimi's pool construction is unrelated. + builder.__class__ = KimiAiterMLAGDNMetadataBuilder + builder._build_gdn_capture_metadata = lambda bs: None + runner.h2d_owner.begin() + md, _ = builder.build_for_cudagraph_capture(3) + saved = md.kv_indptr.clone() + assert ( + runner.forward_vars["kv_indptr"]._publication._epoch == runner.h2d_owner.epoch + ) + assert_repeat_does_not_write(runner, lambda: builder.build_for_cudagraph_capture(2)) + finish(runner) + torch.cuda.synchronize() + assert saved.tolist() == [0, 1, 2, 3] + + +def make_qwen_runner(monkeypatch, pp=1, *, bind=True): + from atom.model_ops.attentions.gdn_attn import GDNAttentionMetadataBuilder + from atom.model_ops.attentions.qwen4_exp_attn import Qwen4ExpMetadataBuilder + + runner, _ = make_runner(monkeypatch, "mha", pp=1, tbo=False, bind=False) + + # Invoke the real Qwen allocation while using the already allocated base. + def init(self, model_runner, **kwargs): + self.model_runner = model_runner + self.device = model_runner.device + self.max_bs = model_runner.max_bs + self.max_num_batched_tokens = model_runner.max_num_batched_tokens + + monkeypatch.setattr(GDNAttentionMetadataBuilder, "__init__", init) + # Bind a fresh runner after all backend fields have been allocated. + builder = Qwen4ExpMetadataBuilder(runner) + runner.attn_metadata_builder = builder + runner.config.pipeline_parallel_size = pp + runner.ple_conv_state = torch.zeros(8, 2, device=runner.device) + runner.ple_ngram_state = torch.zeros(8, 2, device=runner.device) + runner._init_forward_vars_ring() + if bind: + runner._init_h2d_publication() + return runner, builder + + +@pytest.mark.parametrize("pp", [1, 2]) +def test_qsa_ple_ragged_padding_delayed_graph_consumer(monkeypatch, pp): + runner, builder = make_qwen_runner(monkeypatch, pp) + graphs, outputs = [], [] + # Isolated consumer graphs per slot; production PP still requires eager. + for var in runner._fv_ring: + out = torch.empty(12, dtype=torch.int64, device=runner.device) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + torch.add( + var["qsa_logical_positions"].gpu[:12], + var["qsa_token_to_req"].gpu[:12], + out=out, + ) + graphs.append(graph) + outputs.append(out) + saved = [] + for step, counts in enumerate(([3, 1, 4], [1], [], [4, 4, 4], [2, 3])): + runner._advance_forward_vars() + runner._gate_staging_reuse() + var = runner.forward_vars + n = sum(counts) + positions = list(range(step, step + n)) + var["positions"].np[:n] = positions + var["slot_mapping"].gpu[:12] = torch.arange(12, device=runner.device) + 64 + state_indices = torch.arange( + len(counts), dtype=torch.int32, device=runner.device + ) + md = SimpleNamespace( + slot_mapping=var["slot_mapping"].gpu[:12], + block_tables=var["block_tables"].gpu, + context_lens=var["context_lens"].gpu, + cu_seqlens_q=var["cu_seqlens_q"].gpu, + max_seqlen_k=128, + gdn_metadata=SimpleNamespace( + spec_state_indices_tensor=None, + non_spec_state_indices_tensor=state_indices, + non_spec_state_indices_in_tensor=state_indices, + num_accepted_tokens=None, + ), + ) + cached = [i % 2 for i in range(len(counts))] + batch = SimpleNamespace(num_cached_tokens=cached) + var["ple_has_initial_state"].gpu.fill_(True) + torch.cuda._sleep(3_000_000) + args = (md, len(counts), 12, np.asarray(counts, dtype=np.int64), n) + qsa = builder._build_qsa_metadata(*args) + ple = builder._build_ple_metadata(batch, md, len(counts), is_prefill=True) + graphs[runner._fv_idx].replay() + ids = [i for i, count in enumerate(counts) for _ in range(count)] + saved.append( + ( + outputs[runner._fv_idx].clone(), + [p + i for p, i in zip(positions, ids)] + [-2] * (12 - n), + ) + ) + saved.append( + ( + qsa.compressed_slot_mapping.clone(), + [ + (64 + i) // 4 if (p + 1) % 4 == 0 else -1 + for i, p in enumerate(positions) + ] + + [-1] * (12 - n), + ) + ) + saved.append((ple.has_initial_state.clone(), [bool(x) for x in cached])) + saved.append( + ( + var["ple_has_initial_state"].gpu[len(counts) :].clone(), + [True] * (4 - len(counts)), + ) + ) + assert_repeat_does_not_write( + runner, lambda args=args: builder._build_qsa_metadata(*args) + ) + assert_repeat_does_not_write( + runner, + lambda batch=batch, md=md, counts=counts: builder._build_ple_metadata( + batch, md, len(counts), is_prefill=True + ), + ) + finish(runner) + torch.cuda.synchronize() + for actual, expected in saved: + assert actual.cpu().tolist() == expected + + +def test_qsa_capture_rejected_before_raw_source_write(monkeypatch): + runner, builder = make_qwen_runner(monkeypatch) + runner.h2d_owner.begin() + before = runner.forward_vars["qsa_token_to_req"].cpu.clone() + # Preflight must fail before a host fill even for an empty publication. + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph), pytest.raises(PublicationError, match="capture"): + builder._build_qsa_metadata(None, 0, 0, np.array([], dtype=np.int64), 0) + assert torch.equal(before, runner.forward_vars["qsa_token_to_req"].cpu) + finish(runner) + + +@pytest.mark.parametrize("pp", [1, 2]) +@pytest.mark.parametrize("cached", [False, True]) +def test_mla_sparse_prefill_prefixes_and_tail(monkeypatch, pp, cached): + from atom.model_ops.attentions import aiter_mla + + runner, builder = make_runner(monkeypatch, pp=pp, sparse=True, tbo=False) + monkeypatch.setattr(aiter_mla, "get_mla_metadata_v1", lambda *a, **kw: None) + saved = [] + for step in range(4): + counts = np.asarray([3, 2] if cached else [10 + step, 3], dtype=np.int32) + prefix = np.asarray([9 + step, 5] if cached else [0, 0], dtype=np.int32) + lens = counts + prefix + tokens = int(counts.sum()) + batch = decode_batch(2) + batch.total_seqs_num_decode = batch.total_tokens_num_decode = 0 + batch.total_seqs_num_prefill = 2 + batch.total_tokens_num_prefill = batch.total_tokens_num = tokens + batch.num_scheduled_tokens = counts + batch.context_lens = lens + batch.num_cached_tokens = prefix.tolist() + batch.last_block_num_tokens = ((lens - 1) % 16 + 1).tolist() + runner._advance_forward_vars() + runner._gate_staging_reuse() + builder.publish_cu_seqlens_q(batch, SimpleNamespace(running_bs=4)) + for name in ( + "cu_seqlen_ks", + "cu_seqlen_ke", + "sparse_kv_indptr", + "kv_last_page_lens", + ): + runner.forward_vars[name].gpu.fill_(-99) + torch.cuda._sleep(3_000_000) + md, _ = builder.prepare_prefill(batch, 4) + starts, ends, bids, selected = [], [], [], [] + base = 0 + for i, (n, old) in enumerate(zip(counts, prefix)): + for j in range(int(n)): + starts.append(base) + ends.append(base + int(old) + j + 1) + bids.append(i) + selected.append(min(int(old) + j + 1, 8)) + base += int(n + old) + saved.extend( + [ + (md.cu_seqlen_ks.clone(), starts), + (md.cu_seqlen_ke.clone(), ends), + (md.batch_id_per_q_token.clone(), bids), + (md.sparse_kv_indptr.clone(), [0] + np.cumsum(selected).tolist()), + ( + runner.forward_vars["cu_seqlen_ke"].gpu[tokens:].clone(), + [-99] * (64 - tokens), + ), + ] + ) + if cached: + saved.append( + ( + runner.forward_vars["kv_last_page_lens"].gpu.clone(), + batch.last_block_num_tokens + [0, 0], + ) + ) + assert_repeat_does_not_write( + runner, lambda batch=batch: builder.prepare_prefill(batch, 4) + ) + finish(runner) + torch.cuda.synchronize() + for actual, expected in saved: + assert actual.cpu().tolist() == expected + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +@pytest.mark.parametrize("producer", ["prefill", "mrope_prefill", "mrope_decode"]) +def test_attention_reentry_preserves_sources_and_first_consumer( + monkeypatch, transport, producer +): + from tests.test_h2d_runner_publication import runner_with_buffers + + if producer == "prefill": + runner, builder = make_runner( + monkeypatch, "mha", tbo=False, transport=transport + ) + + def produce(changed): + batch = decode_batch(2) + batch.total_seqs_num_decode = batch.total_tokens_num_decode = 0 + batch.total_seqs_num_prefill = 2 + batch.total_tokens_num_prefill = batch.total_tokens_num = 5 + batch.num_scheduled_tokens = np.array([3, 2], dtype=np.int32) + batch.context_lens = np.array([12, 7], dtype=np.int32) + changed * np.array( + [1, -1], dtype=np.int32 + ) + batch.num_cached_tokens = ( + batch.context_lens - batch.num_scheduled_tokens + ).tolist() + return builder.prepare_prefill(batch, 4) + + names = ("cu_seqlens_k", "context_lens", "positions", "slot_mapping") + cu = runner.forward_vars["cu_seqlens_q"] + cu.cpu[:5] = torch.tensor([0, 3, 5, 5, 5]) + cu.gpu.copy_(cu.cpu) + else: + runner = runner_with_buffers(monkeypatch, transport) + builder = runner.attn_metadata_builder + + def produce(changed): + ends = np.array([12, 22], dtype=np.int32) + changed + batch = SimpleNamespace( + total_tokens_num_decode=4, + total_tokens_num_prefill=4, + req_ids=[10, 20], + context_lens=ends, + num_cached_tokens=ends - 2, + mrope_positions_by_req={}, + mrope_position_deltas={}, + ) + if producer == "mrope_prefill": + return builder._build_mrope_prefill_positions(batch) + return builder._build_mrope_decode_positions( + batch, ends, 2, running_tokens=8 + ) + + names = ("mrope_positions",) + + runner._gate_staging_reuse() + produce(0) + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + sources = [runner.forward_vars[name] for name in names] + expected = [buf.gpu.clone() for buf in sources] + observed = [torch.empty_like(value) for value in expected] + runner._gate_staging_reuse() + torch.cuda._sleep(20_000_000) + produce(0) + host_before = [buf.cpu.clone() for buf in sources] + for out, buf in zip(observed, sources): + out.copy_(buf.gpu) + try: + with pytest.raises(PublicationError, match="republish_reason"): + produce(-1) + finally: + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + for name, buf, before, actual, reference in zip( + names, sources, host_before, observed, expected + ): + assert torch.equal(buf.cpu, before), name + assert torch.equal(actual, reference), name + + +@pytest.mark.parametrize("kind", ["mha", "mla"]) +@pytest.mark.parametrize("transport", ["direct", "packed"]) +@pytest.mark.parametrize("pp", [1, 2]) +def test_shared_maps_skip_h2d_per_slot_and_tbo_buffer(monkeypatch, kind, transport, pp): + from atom.model_engine.sequence import new_block_table + + runner, builder = make_runner( + monkeypatch, kind, pp=pp, tbo=True, transport=transport + ) + rows = [new_block_table(row) for row in decode_batch(3).block_tables] + published = {} + observed = [] + for step in range(9): + if step == 4: + rows = list(reversed(rows)) + if step == 7: + rows[0].append(31) + batch = decode_batch(3, step=step % 4) + batch.block_tables = rows + start_decode(runner, builder, batch, 1) + torch.cuda._sleep(2_000_000) + builder.prepare_decode(batch, 4, 4, 1) + for name, selected, width in ( + ("block_tables", rows, 4), + ("ub0_block_tables", rows[:2], 2), + ("ub1_block_tables", rows[2:], 2), + ): + buf = runner.forward_vars[name] + mapping = tuple(tuple(row) for row in selected) + key = (runner._fv_idx, name) + changed = published.get(key) != mapping + assert (buf._publication._epoch == runner.h2d_owner.epoch) == changed + published[key] = mapping + expected = torch.zeros_like(buf.cpu[:width]) + for i, row in enumerate(selected): + expected[i, : len(row)] = torch.tensor(list(row), dtype=torch.int32) + observed.append((buf.gpu[:width].clone(), expected)) + finish(runner) + torch.cuda.synchronize() + for actual, expected in observed: + assert torch.equal(actual.cpu(), expected) diff --git a/tests/test_h2d_draft_publication.py b/tests/test_h2d_draft_publication.py new file mode 100644 index 0000000000..56e79c3e5b --- /dev/null +++ b/tests/test_h2d_draft_publication.py @@ -0,0 +1,314 @@ +# SPDX-License-Identifier: MIT +"""Draft device staging and GDN state indices through real forward slots.""" + +import os +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from atom.utils import CpuGpuBuffer +from atom.utils.h2d import PublicationError + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1", reason="set RUN_H2D_GPU_TESTS=1" +) + + +def make_runner( + monkeypatch, *, pp_size=1, num_spec=3, replay=False, base=None, max_bs=4 +): + from atom.model_engine.model_runner import ModelRunner + from atom.model_ops.attentions.gdn_attn import GDNStateMixin + + monkeypatch.setenv("ATOM_H2D_BACKEND", "direct") + monkeypatch.setenv("ATOM_ENABLE_REPLAYSSM", str(int(replay))) + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace( + pipeline_parallel_size=pp_size, hf_config=SimpleNamespace() + ) + runner.enforce_eager = True + runner.tokenID_processor = SimpleNamespace(num_bonus=None) + runner.drafter = SimpleNamespace(mtp_k=num_spec, runner=runner) + runner.forward_vars = { + name: CpuGpuBuffer(32, dtype=torch.int32, device=runner.device) + for name in ("input_ids", "decode_src") + } + runner.forward_vars["draft_next_tokens"] = CpuGpuBuffer( + 4, dtype=torch.int32, device=runner.device, publication_group="draft_anchors" + ) + runner.replayssm_write_pos = torch.full( + (64,), -1, dtype=torch.int32, device=runner.device + ) + cls = base or GDNStateMixin + builder = cls.__new__(cls) + builder.model_runner = runner + builder.device = runner.device + builder.max_bs = max_bs + builder._init_gdn_state(runner) + runner.attn_metadata_builder = builder + runner._init_forward_vars_ring() + runner._init_h2d_publication() + return runner, builder + + +def batch_and_metadata(builder, step, count, *, prefill=False): + from atom.utils.forward_context import AttentionMetaData + + width = builder.num_spec + 1 if builder.use_spec_decode and not prefill else 1 + starts = [step + 3 * i + 1 for i in range(count)] + slots = [ + [s] if builder.replayssm or prefill else [s + 7 * j for j in range(width)] + for s in starts + ] + forks = [s + 25 if i % 2 else -1 for i, s in enumerate(starts)] if prefill else [] + batch = SimpleNamespace( + state_slots=slots, + state_fork_srcs=forks, + total_seqs_num=count, + total_seqs_num_decode=0 if prefill else count, + total_seqs_num_prefill=count if prefill else 0, + total_tokens_num=count * width, + total_tokens_num_decode=0 if prefill else count * width, + total_tokens_num_prefill=count * width if prefill else 0, + ) + md = AttentionMetaData( + cu_seqlens_q=torch.arange(count + 1, dtype=torch.int32, device=builder.device) + * width + ) + return batch, md + + +@pytest.mark.parametrize("pp_size", [1, 2]) +@pytest.mark.parametrize("num_spec,replay", [(0, False), (3, False), (3, True)]) +def test_gdn_decode_indices_select_slot_and_preserve_padding( + monkeypatch, pp_size, num_spec, replay +): + runner, builder = make_runner( + monkeypatch, pp_size=pp_size, num_spec=num_spec, replay=replay + ) + saved = [] + for step, count in enumerate((3, 1, 4, 0, 2, 3)): + runner._advance_forward_vars() + runner._gate_staging_reuse() + batch, md = batch_and_metadata(builder, step, count) + builder._attach_gdn_decode_metadata(batch, md, prepare_block_tables=False) + attrs = ( + [("spec_state_indices_tensor", "spec_state_indices")] + if num_spec + else [ + ("non_spec_state_indices_tensor", "non_spec_state_indices"), + ("non_spec_state_indices_in_tensor", "non_spec_state_indices_in"), + ] + ) + if replay: + attrs.append(("slot_idx", "non_spec_state_indices")) + for attr, name in attrs: + buf = runner.forward_vars[name] + assert ( + getattr(md.gdn_metadata, attr).untyped_storage().data_ptr() + == buf.gpu.untyped_storage().data_ptr() + ) + assert buf._publication._epoch == runner.h2d_owner.epoch + if name == "spec_state_indices": + expected = np.zeros((count, num_spec + 1), dtype=np.int32) + for row, slots in enumerate(batch.state_slots): + expected[row, : len(slots)] = slots + expected = expected.tolist() + [[-1] * (num_spec + 1)] * (4 - count) + else: + expected = [row[0] for row in batch.state_slots] + [-1] * (4 - count) + saved.append((buf.gpu.clone(), expected)) + before = { + member.name: member.source.clone() + for member in runner.h2d_groups["gdn_state"].members + } + with pytest.raises(PublicationError, match="republish_reason"): + builder.prepare_state_indices(batch, with_spec=bool(num_spec)) + assert all( + torch.equal(runner.forward_vars[name].cpu, value) + for name, value in before.items() + ) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for actual, expected in saved: + assert actual.cpu().tolist() == expected + + +@pytest.mark.parametrize("pp_size", [1, 2]) +@pytest.mark.parametrize("num_spec,replay", [(0, False), (3, False), (3, True)]) +def test_gdn_ragged_decode_reuses_buffers_and_restores_graph_padding( + monkeypatch, pp_size, num_spec, replay +): + runner, builder = make_runner( + monkeypatch, pp_size=pp_size, num_spec=num_spec, replay=replay + ) + prefix = ( + builder.spec_query_start_loc if num_spec else builder.non_spec_query_start_loc + ) + # This graph is only a reader of metadata; PP model execution remains eager. + # It detects stale padding and addresses after each preparation. + outputs = [torch.empty_like(prefix)] + sources = [prefix] + if num_spec: + sources += [ + builder.spec_sequence_masks, + builder.spec_token_indx, + builder.num_accepted_tokens, + ] + outputs += [torch.empty_like(x) for x in sources[1:]] + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for dst, src in zip(outputs, sources): + dst.copy_(src) + saved = [] + for step, lengths in enumerate(([4, 1, 2], [1], [2, 4, 1, 3], [], [1, 3])): + if not num_spec: + lengths = [1] * len(lengths) + count, total = len(lengths), sum(lengths) + runner._advance_forward_vars() + runner._gate_staging_reuse() + batch, md = batch_and_metadata(builder, step, count) + batch.total_tokens_num = batch.total_tokens_num_decode = total + cu = [0, *np.cumsum(lengths).tolist()] + # Both scheduled and padded ends must be the actual ragged total. + md.cu_seqlens_q = torch.tensor( + cu + [total] * (4 - count), dtype=torch.int32, device=runner.device + ) + bonuses = [i % 4 for i in range(count)] + runner.tokenID_processor.num_bonus = bonuses if num_spec else None + torch.cuda._sleep(5_000_000) + builder._attach_gdn_decode_metadata(batch, md, prepare_block_tables=False) + if num_spec: + assert md.gdn_metadata.spec_token_indx.numel() == total + assert md.gdn_metadata.non_spec_token_indx.numel() == 0 + assert ( + md.gdn_metadata.spec_token_indx.untyped_storage().data_ptr() + == builder.spec_token_indx.data_ptr() + ) + graph.replay() + saved.append(([x.clone() for x in outputs], cu, count, bonuses)) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for actual, cu, count, bonuses in saved: + assert actual[0].cpu().tolist() == cu + [cu[-1]] * (4 - count) + if num_spec: + assert actual[1].cpu().tolist() == [True] * count + [False] * (4 - count) + assert actual[2].cpu().tolist() == list(range(16)) + assert actual[3].cpu().tolist() == [b + 1 for b in bonuses] + [1] * ( + 4 - count + ) + + +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_gdn_fork_indices_delayed_reuse_and_draft_graph(monkeypatch, pp_size): + runner, builder = make_runner(monkeypatch, pp_size=pp_size) + bank = torch.arange(64, dtype=torch.int32, device=runner.device) * 11 + graphs, outputs = [], [] + for variables in runner._fv_ring: + out = torch.empty(4, dtype=torch.int32, device=runner.device) + indices = variables["non_spec_state_indices_in"].gpu + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + torch.index_select(bank, 0, indices, out=out) + graphs.append(graph) + outputs.append(out) + saved = [] + for step in range(6): + runner._advance_forward_vars() + runner._gate_staging_reuse() + batch, _ = batch_and_metadata(builder, step, 4, prefill=True) + torch.cuda._sleep(5_000_000) + builder.prepare_state_indices(batch) + builder.non_spec_state_indices_tensor.copy_to_gpu(4) + builder.non_spec_state_indices_in_tensor.copy_to_gpu(4) + graphs[runner._fv_idx].replay() + expected = [ + src if src >= 0 else row[0] + for src, row in zip(batch.state_fork_srcs, batch.state_slots) + ] + saved.append((outputs[runner._fv_idx].clone(), [i * 11 for i in expected])) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for actual, expected in saved: + assert actual.cpu().tolist() == expected + + +@pytest.mark.parametrize("replay", [False, True]) +def test_gdn_prefill_fork_metadata_uses_current_slot(monkeypatch, replay): + runner, builder = make_runner(monkeypatch, pp_size=2, replay=replay) + runner._advance_forward_vars() + runner._gate_staging_reuse() + batch, md = batch_and_metadata(builder, 2, 3, prefill=True) + actual = builder.prepare_gdn_metadata( + batch, md, is_prefill=True, prepare_block_tables=False + ) + expected_out = [row[0] for row in batch.state_slots] + expected_in = [ + src if src >= 0 else row[0] + for src, row in zip(batch.state_fork_srcs, batch.state_slots) + ] + assert ( + actual.non_spec_state_indices_tensor.data_ptr() + == runner.forward_vars["non_spec_state_indices"].gpu.data_ptr() + ) + assert ( + actual.non_spec_state_indices_in_tensor.data_ptr() + == runner.forward_vars["non_spec_state_indices_in"].gpu.data_ptr() + ) + runner._mark_staging_h2d_enqueued() + assert actual.non_spec_state_indices_tensor.cpu().tolist() == expected_out + assert actual.non_spec_state_indices_in_tensor.cpu().tolist() == expected_in + assert not actual.has_initial_state.any() + runner._record_forward_vars_event() + + +@pytest.mark.parametrize("replay", [False, True]) +def test_capture_preparation_uses_checked_state_sources(monkeypatch, replay): + runner, builder = make_runner(monkeypatch, replay=replay) + for count in (4, 1, 3): + runner.h2d_owner.begin() + builder._prepare_state_indices_for_capture(count) + md = builder._build_gdn_capture_metadata(count) + before = builder.spec_state_indices_tensor.cpu.clone() + with pytest.raises(PublicationError, match="republish_reason"): + builder._prepare_state_indices_for_capture(count) + assert torch.equal(builder.spec_state_indices_tensor.cpu, before) + runner.h2d_owner.finish() + expected = np.arange(count * 4).reshape(count, 4) + if replay: + expected = np.repeat(np.arange(count)[:, None], 4, axis=1) + np.testing.assert_array_equal( + md.spec_state_indices_tensor.cpu().numpy(), expected + ) + runner.h2d_owner.begin() + before = builder.non_spec_state_indices_tensor.cpu.clone() + with ( + torch.cuda.graph(torch.cuda.CUDAGraph()), + pytest.raises(PublicationError, match="before actual graph capture"), + ): + builder._prepare_state_indices_for_capture(2) + assert torch.equal(builder.non_spec_state_indices_tensor.cpu, before) + runner.h2d_owner.finish() + + +def test_anchor_duplicate_rejected_before_source_write(monkeypatch): + from atom.spec_decode.drafter import Drafter + + runner, _ = make_runner(monkeypatch) + runner._gate_staging_reuse() + torch.cuda._sleep(5_000_000) + first = Drafter.anchors_to_gpu(runner.drafter, [4, -1, 9]).clone() + before = runner.forward_vars["draft_next_tokens"].cpu.clone() + with pytest.raises(PublicationError, match="republish_reason"): + Drafter.anchors_to_gpu(runner.drafter, [7, 8, 9]) + assert torch.equal(runner.forward_vars["draft_next_tokens"].cpu, before) + runner._mark_staging_h2d_enqueued() + assert first.cpu().tolist() == [4, -1, 9] + with pytest.raises(PublicationError, match="sealed"): + Drafter.anchors_to_gpu(runner.drafter, [1]) + assert torch.equal(runner.forward_vars["draft_next_tokens"].cpu, before) diff --git a/tests/test_h2d_publication.py b/tests/test_h2d_publication.py new file mode 100644 index 0000000000..6209bee329 --- /dev/null +++ b/tests/test_h2d_publication.py @@ -0,0 +1,678 @@ +# SPDX-License-Identifier: MIT +"""Publication contracts exercised with real tensors, including delayed GPUs.""" + +import gc +import os +import sys +import weakref +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace + +import pytest +import torch + +from atom.utils import CpuGpuBuffer +from atom.utils.h2d import PublicationError, PublicationOwner, PublicationRegistry + + +def buffer(*shape, dtype=torch.int32, device="cpu"): + return CpuGpuBuffer( + *shape, + dtype=dtype, + device=device, + pin_memory=device != "cpu", + with_numpy=dtype != torch.bfloat16, + ) + + +def setup(device="cpu", shape=(16,), dtype=torch.int32, unit="rows"): + owner = PublicationOwner(device, torch.cuda.Event() if device != "cpu" else None) + a, b = buffer(*shape, dtype=dtype, device=device), buffer( + *shape, dtype=dtype, device=device + ) + x, y = owner.bind(a, "a", unit=unit), owner.bind(b, "b", unit=unit) + group = owner.group("pair", (x, y)) + return owner, a, b, x, y, group + + +@pytest.mark.parametrize("bad", [-1, 17, 1.5, True]) +def test_group_validation_is_atomic_and_retryable(bad): + owner, a, b, _, _, group = setup() + owner.begin() + a.cpu.fill_(11) + b.cpu.fill_(22) + group.set_count(a, 8) + group.set_count(b, bad) + with pytest.raises((TypeError, ValueError)): + group.publish(group.counts) + assert torch.all(a.gpu == 0) and torch.all(b.gpu == 0) + # Correct one producer without losing the other producer's staged prefix. + group.set_count(b, 8) + group.publish(group.counts) + assert torch.all(a.gpu[:8] == 11) and torch.all(a.gpu[8:] == 0) + assert torch.all(b.gpu[:8] == 22) + + +def test_direct_and_groups_share_epoch_including_later_phases(): + owner, a, b, x, _, group = setup() + other = owner.group("other", (x,)) + owner.begin() + a.copy_to_gpu(8) + with pytest.raises(PublicationError, match="republish_reason"): + group.publish((8, 8)) + b.copy_to_gpu(8) # Failed group did not consume b's publication. + owner.finish() + owner.resume() + with pytest.raises(PublicationError, match="republish_reason"): + other.publish((8,)) + with pytest.raises(PublicationError, match="cannot begin"): + owner.begin() + + +def test_another_thread_cannot_publish_or_resume_the_owner(): + owner, a, _, _, _, _ = setup() + owner.begin() + with ThreadPoolExecutor(max_workers=1) as worker: + with pytest.raises(PublicationError, match="owner thread"): + worker.submit(a.copy_to_gpu).result() + a.copy_to_gpu() + owner.finish() + with pytest.raises(PublicationError, match="owner thread"): + worker.submit(owner.resume).result() + owner.resume() + with pytest.raises(PublicationError, match="republish_reason"): + a.copy_to_gpu() + + +def test_omitted_members_differ_from_explicit_empty_publication(): + owner, a, b, _, _, group = setup() + owner.begin() + group.publish((None, None)) + group.publish((0, None)) + with pytest.raises(PublicationError, match="republish_reason"): + a.copy_to_gpu(0) + b.cpu.fill_(22) + b.copy_to_gpu(8) + b.gpu.fill_(-88) + group.publish((None, None)) + assert torch.all(b.gpu == -88) + + +def test_reason_does_not_release_source_and_reacquire_does_not_reset_ledger(): + owner, a, _, x, _, _ = setup() + owner.begin() + a.cpu.fill_(11) + a.copy_to_gpu() + first = a.gpu.clone() + with pytest.raises(PublicationError, match="acquire_write"): + a.copy_to_gpu(republish_reason="publish postprocess correction") + with pytest.raises(PublicationError, match="republish_reason"): + x.acquire_write() + x.acquire_write(republish_reason="publish postprocess correction").fill_(22) + with pytest.raises(PublicationError, match="republish_reason"): + a.copy_to_gpu() + a.copy_to_gpu(republish_reason="publish postprocess correction") + assert torch.all(first == 11) and torch.all(a.gpu == 22) + + +def test_aliases_rejected_across_owners_but_disjoint_regions_allowed(): + registry = PublicationRegistry() + owner = PublicationOwner("cpu", registry=registry) + other = PublicationOwner("cpu", registry=registry) + a, b = buffer(16), buffer(8) + a.cpu, a.gpu = a.cpu[:8], a.gpu[:8] + owner.bind(a, "a") + b.gpu = a.gpu[2:6] + b.cpu = b.cpu[:4] + with pytest.raises(ValueError, match="destination overlaps"): + other.bind(b, "alias") + c, d = buffer(16), buffer(16) + d.cpu, d.gpu = c.cpu[8:], c.gpu[8:] + c.cpu, c.gpu = c.cpu[:8], c.gpu[:8] + owner.bind(c, "left") + other.bind(d, "right") + + +def test_shared_host_arena_cannot_be_registered_with_independent_reuse_state(): + registry = PublicationRegistry() + one, two = PublicationOwner("cpu", registry=registry), PublicationOwner( + "cpu", registry=registry + ) + a, b = buffer(16), buffer(16) + b.cpu = a.cpu + one.bind(a, "a") + with pytest.raises(ValueError, match="source overlaps"): + two.bind(b, "b") + + +def test_duplicate_destination_in_group_rejected(): + owner, _, _, x, _, _ = setup() + with pytest.raises(ValueError, match="destination twice"): + owner.group("duplicate", (x, x)) + + +@pytest.mark.parametrize( + "dtype", [torch.int32, torch.int64, torch.float32, torch.bfloat16, torch.bool] +) +def test_mixed_representation_is_bitwise_and_tail_preserved(dtype): + owner, a, _, _, _, group = setup(shape=(7,), dtype=dtype, unit="bytes") + src = a.cpu.view(torch.uint8) + dst = a.gpu.view(torch.uint8) + src.copy_(torch.arange(src.numel(), dtype=torch.uint8) * 31) + dst.fill_(197) + owner.begin() + group.publish((src.numel() - 1, None)) + assert torch.equal(src[:-1], dst[:-1]) + assert dst[-1] == 197 + + +def test_rows_and_flat_prefix_are_distinct_and_legacy_copy_keeps_rows(): + owner, a, b, _, _, group = setup(shape=(3, 16), unit="elements") + a.cpu.fill_(11) + b.cpu.fill_(22) + owner.begin() + group.publish((3 * 5, None)) + b.copy_to_gpu(1) + assert torch.all(a.gpu.flatten()[:15] == 11) + assert torch.all(a.gpu.flatten()[15:] == 0) + assert torch.all(b.gpu[0] == 22) and torch.all(b.gpu[1:] == 0) + + +def test_noncontiguous_direct_does_not_copy_unselected_columns(): + owner = PublicationOwner("cpu") + a = buffer(3, 8) + backing = a.gpu + a.cpu, a.gpu = a.cpu[:, ::2], a.gpu[:, ::2] + a.cpu.fill_(11) + owner.bind(a, "strided") + owner.begin() + a.copy_to_gpu(2) + assert torch.all(backing[:2, ::2] == 11) + assert torch.all(backing[:, 1::2] == 0) and torch.all(backing[2] == 0) + + +def test_partial_enqueue_poison_prevents_next_forward(monkeypatch): + owner, a, b, _, y, group = setup() + owner.begin() + a.cpu.fill_(11) + b.cpu.fill_(22) + + def fail(count): + raise RuntimeError("injected enqueue failure") + + monkeypatch.setattr(y, "_copy", fail) + with pytest.raises(RuntimeError, match="injected"): + group.publish((8, 8)) + assert torch.all(a.gpu[:8] == 11) and torch.all(b.gpu == 0) + owner.drain() + with pytest.raises(PublicationError, match="failed"): + owner.begin() + with pytest.raises(PublicationError, match="failed"): + group.publish((None, None)) + + +@pytest.mark.parametrize("failure_phase", ["begin", "finish", "acquire"]) +def test_completion_failure_poison_prevents_source_reuse(failure_phase): + class Completion: + fail = False + + def synchronize(self): + if self.fail: + raise RuntimeError("event failed") + + def record(self, stream): + if self.fail: + raise RuntimeError("event failed") + + owner, a, _, binding, _, group = setup() + event = owner.completion = Completion() + if failure_phase != "begin": + owner.begin() + a.copy_to_gpu() + event.fail = True + with pytest.raises(RuntimeError, match="event failed"): + if failure_phase == "begin": + owner.begin() + elif failure_phase == "finish": + owner.finish() + else: + binding.acquire_write(republish_reason="update after first consumer") + owner.drain() + with pytest.raises(PublicationError, match="failed"): + owner.begin() + with pytest.raises(PublicationError, match="failed"): + group.publish((None, None)) + + +def test_strong_reference_and_binding_replacement(): + owner, a, _, _, _, group = setup() + ref = weakref.ref(a.cpu) + owner.begin() + with pytest.raises(PublicationError, match="bound metadata storage is fixed"): + a.cpu = torch.ones_like(a.cpu) + gc.collect() + assert ref() is not None + group.publish((8, None)) + + +@pytest.mark.parametrize("transport", ["packed"]) +def test_transport_selection_preserves_fallback_and_owner_lifecycle( + monkeypatch, transport +): + owner, a, _, _, _, group = setup() + # CPU-only users must not import the Triton backend, even when requested. + monkeypatch.setitem(sys.modules, "atom.utils.packed_h2d", None) + assert group.use_transport(transport) == "direct" + reason = group.fallback_reason + assert reason + with pytest.raises(ValueError, match="H2D transport"): + group.use_transport("unknown") + assert group.fallback_reason == reason + assert group.use_transport("direct") == "direct" + assert group.fallback_reason is None and group._backend is None + owner.begin() + a.cpu.fill_(31) + group.publish((8, None)) + with pytest.raises(PublicationError, match="during initialization"): + group.use_transport(transport) + with pytest.raises(PublicationError, match="republish_reason"): + group.publish((8, None)) + assert torch.all(a.gpu[:8] == 31) and torch.all(a.gpu[8:] == 0) + + +@pytest.mark.parametrize("compact", [False, True]) +def test_v41_reuses_early_query_prefix_or_explicitly_compacts(compact): + from atom.model_ops.attentions.deepseek_v41.metadata import ( + RequestSpan, + prepare_batch_step, + ) + + owner = PublicationOwner("cpu") + buffers = { + "positions": buffer(8), + "cu_seqlens_q": buffer(4), + "batch_id_per_q_token": buffer(8), + "block_tables": buffer(3, 2), + } + cu = buffers["cu_seqlens_q"] + owner.bind(cu, "cu_seqlens_q") + owner.begin() + original = [0, 2, 2, 5] if compact else [0, 2, 5, 5] + cu.cpu.copy_(torch.tensor(original, dtype=torch.int32)) + cu.copy_to_gpu() + first_consumer = cu.gpu.clone() + step = prepare_batch_step( + [RequestSpan(1, 0, 0, 2, 0), RequestSpan(2, 0, 2, 3, 1)], + "cpu", + block_tables=[(0,), (1,)], + buffers=buffers, + running_bs=3, + running_tokens=8, + query_prefix_ready=not compact, + query_prefix_republish_reason=( + "compact zero-token scheduler rows" if compact else None + ), + ) + assert first_consumer.tolist() == original + assert step.cu_seqlens_q.tolist() == [0, 2, 5, 5] + assert buffers["batch_id_per_q_token"].gpu.tolist() == [0, 0, 1, 1, 1, -1, -1, -1] + with pytest.raises(PublicationError, match="republish_reason"): + cu.copy_to_gpu() + + +gpu = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1" or not torch.cuda.is_available(), + reason="set RUN_H2D_GPU_TESTS=1 with a GPU", +) + + +@pytest.mark.parametrize("device", ["cpu", pytest.param("cuda", marks=gpu)]) +@pytest.mark.parametrize("unit", ["rows", "elements", "bytes"]) +def test_repeated_and_changing_prefixes_preserve_contents_and_tail(device, unit): + owner, a, _, binding, _, _ = setup(device, shape=(7, 3), unit=unit) + expected = torch.zeros_like(a.cpu) + width = {"rows": 1, "elements": 3, "bytes": 12}[unit] + for step, rows in enumerate((1, 2, 3, 4, 5, 6, 2, 2, 0, 7), 1): + owner.begin() + a.cpu.fill_(step) + binding._single.publish((rows * width,)) + owner.finish() + if device != "cpu": + owner.completion.synchronize() + expected[:rows] = step + assert torch.equal(a.gpu.cpu(), expected) + + +def test_single_member_failure_poison_and_invalid_counts_are_retryable(monkeypatch): + owner, a, _, binding, _, _ = setup() + owner.begin() + for counts, reasons in (((), None), ((8,), ()), ((True,), None)): + with pytest.raises((ValueError, TypeError)): + binding._single.publish(counts, republish_reasons=reasons) + assert binding._epoch != owner.epoch and torch.all(a.gpu == 0) + + def fail(count): + raise RuntimeError("injected single enqueue failure") + + monkeypatch.setattr(binding, "_copy", fail) + with pytest.raises(RuntimeError, match="single enqueue"): + a.copy_to_gpu(8) + with pytest.raises(PublicationError, match="failed"): + binding._single.publish((None,)) + + +@gpu +def test_delayed_gpu_different_buffers_and_explicit_republication(): + owner, a, b, x, _, group = setup("cuda") + torch.cuda.synchronize() + owner.begin() + torch.cuda._sleep(20_000_000) + a.cpu.fill_(11) + group.publish((16, None)) + first = a.gpu.clone() + b.cpu.fill_(22) + group.publish((None, 16)) + second = b.gpu.clone() + x.acquire_write(republish_reason="postprocess correction").fill_(33) + a.copy_to_gpu(republish_reason="postprocess correction") + owner.finish() + owner.completion.synchronize() + assert torch.all(first.cpu() == 11) + assert torch.all(second.cpu() == 22) + assert torch.all(a.gpu.cpu() == 33) + + +@gpu +def test_slot_rotation_waits_before_payload_changes(): + slots = [setup("cuda") for _ in range(2)] + torch.cuda.synchronize() + observations = [] + for epoch in range(8): + owner, a, _, _, _, group = slots[epoch % 2] + owner.begin() + a.cpu.fill_(epoch) + torch.cuda._sleep(2_000_000) + group.publish((16, None)) + observations.append(a.gpu.clone()) + owner.finish() + torch.cuda.synchronize() + for epoch, result in enumerate(observations): + assert torch.all(result.cpu() == epoch) + + +@gpu +def test_publish_before_graph_replay_and_reject_actual_capture(): + owner, a, _, _, _, group = setup("cuda") + owner.begin() + group.publish((16, None)) + output = torch.empty_like(a.gpu) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + output.copy_(a.gpu) + with pytest.raises(PublicationError, match="actual graph capture"): + group.publish((None, None)) + owner.finish() + pointer = a.gpu.data_ptr() + for value, count in [(11, 16), (22, 8), (33, 0)]: + owner.begin() + a.cpu.fill_(value) + group.publish((count, None)) + owner.finish() + graph.replay() + torch.cuda.synchronize() + assert a.gpu.data_ptr() == pointer + assert torch.all(output[:count].cpu() == value) + assert torch.all(output[:8].cpu() == 22) + assert torch.all(output[8:].cpu() == 11) + + +@gpu +def test_wrong_stream_is_rejected_before_any_submission(): + owner, _, _, _, _, group = setup("cuda") + owner.begin() + with ( + torch.cuda.stream(torch.cuda.Stream()), + pytest.raises(PublicationError, match="compute stream"), + ): + group.publish((16, 16)) + group.publish((16, 16)) + owner.finish() + + +@gpu +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_real_runner_and_prefill_builder_share_reuse_and_publication(pp_size): + from atom.model_engine.model_runner import ModelRunner + from atom.model_ops.attentions.backends import CommonAttentionBuilder + + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace(pipeline_parallel_size=pp_size) + runner.enforce_eager = True + runner.forward_vars = { + "input_ids": buffer(16, device="cuda"), + "decode_src": buffer(16, device="cuda"), + } + for name, shape in ( + ("slot_mapping", (16,)), + ("context_lens", (4,)), + ("block_tables", (4, 8)), + ("cu_seqlens_k", (5,)), + ("num_cached_tokens", (4,)), + ("seq_starts", (4,)), + ): + value = buffer(*shape, device="cuda") + value.publication_group = "prefill" + runner.forward_vars[name] = value + cu = buffer(5, device="cuda") + cu.publication_group = "early" + runner.forward_vars["cu_seqlens_q"] = cu + runner.tokenID_processor = SimpleNamespace( + input_ids=runner.forward_vars["input_ids"] + ) + runner._init_forward_vars_ring() + runner._init_h2d_publication() + + # Avoid constructing a model, but exercise actual production methods and + # tensors rather than checking their spelling or number of copy calls. + class Builder(CommonAttentionBuilder): + __abstractmethods__ = frozenset() + + Builder.__abstractmethods__ = frozenset() + builder = Builder.__new__(Builder) + builder.model_runner = runner + results = [] + for iteration in range(6): + runner._advance_forward_vars() + runner._gate_staging_reuse() + var = runner.forward_vars + for value in var.values(): + value.cpu.fill_(iteration + 1) + var["cu_seqlens_q"].copy_to_gpu(3) + torch.cuda._sleep(2_000_000) + ctx = builder._upload_prefill_mirrors( + 2, 4, 8, iteration % 2 == 0, var["context_lens"].np[:2] + ) + assert ("block_tables" in ctx) == (iteration % 2 == 0) + assert ctx["context_lens"].shape == (4,) + results.append(ctx["context_lens"].clone()) + with pytest.raises(PublicationError, match="republish_reason"): + var["context_lens"].copy_to_gpu(4) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for iteration, result in enumerate(results): + assert torch.all(result.cpu() == iteration + 1) + + +@gpu +@pytest.mark.parametrize("prefill", [False, True]) +def test_rapidserve_entries_gate_sources_on_the_selected_stream(prefill): + from atom.model_engine.model_runner import RapidServeModelRunner + + runner = RapidServeModelRunner.__new__(RapidServeModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace(pipeline_parallel_size=2) + runner.enforce_eager = True + data = buffer(16, device="cuda") + data.publication_group = "early" + runner.forward_vars = {"input_ids": data, "decode_src": buffer(16, device="cuda")} + runner.tokenID_processor = SimpleNamespace(input_ids=data) + runner._init_forward_vars_ring() + runner._init_h2d_publication() + streams = {i: torch.cuda.Stream() for i in range(2)} + runner._decode_streams = runner._prefill_streams = streams + runner._done_event = torch.cuda.Event() + runner._model_fwd_event = torch.cuda.Event() + + def prepare(batch): + buf = runner.forward_vars["input_ids"] + buf.cpu.fill_(batch.value) + torch.cuda._sleep(2_000_000) + return buf.copy_to_gpu(), None, None, None, True, False + + runner.prepare_model = prepare + runner.run_model = lambda inputs, batch: (inputs.clone(), None) + runner.postprocess = lambda batch, logits, *args, **kwargs: logits.clone() + runner.sampler = lambda logits, *args: logits + runner._record_kv_cache_ready = lambda batch: None + results = [] + for value in range(6): + batch = SimpleNamespace(cu_stream_fraction=value % 2, value=value + 1) + output = runner.prefill_forward(batch) if prefill else runner.forward(batch) + results.append(output) + torch.cuda.synchronize() + for i, output in enumerate(results): + assert (output if prefill else output.cpu().tolist()) == [i + 1] * 16 + + +def stream_setup(fast, backend): + owner = PublicationOwner("cuda", torch.cuda.Event()) + if fast: + if owner._get_current_stream is None: + pytest.skip("current-stream tuple API unavailable") + else: + owner._get_current_stream = None + buffers = [CpuGpuBuffer(8, dtype=torch.int32, device="cuda") for _ in range(2)] + members = [owner.bind(buf, f"value_{i}") for i, buf in enumerate(buffers)] + group = owner.group("pair", members) + if backend == "packed": + assert group.use_transport("packed") == "packed" + return owner, buffers, group + + +@gpu +@pytest.mark.parametrize("fast", [False, True], ids=["public", "tuple"]) +@pytest.mark.parametrize("backend", ["direct", "packed"]) +def test_stream_alias_and_new_epoch_preserve_values_and_ledger(fast, backend): + owner, buffers, group = stream_setup(fast, backend) + first, second = torch.cuda.Stream(), torch.cuda.Stream() + # Another Python wrapper with the same PyTorch stream ID is valid. + alias = torch.cuda.Stream( + stream_id=first.stream_id, + device_index=first.device_index, + device_type=first.device_type, + ) + external = torch.cuda.ExternalStream(first.cuda_stream, device=first.device) + torch.cuda.synchronize() # Complete constructor writes before crossing streams. + with torch.cuda.stream(first): + owner.begin() + group.publish((8, 8)) # Compile packing before the delayed queue. + owner.finish() + owner.begin() + for i, buf in enumerate(buffers): + buf.cpu.fill_(11 + i) + with ( + torch.cuda.stream(second), + pytest.raises(PublicationError, match="compute stream"), + ): + group.publish((8, 8)) + # ExternalStream may have a distinct logical ID despite the same HIP handle. + if external != first: + with ( + torch.cuda.stream(external), + pytest.raises(PublicationError, match="compute stream"), + ): + group.publish((8, 8)) + else: + with torch.cuda.stream(external): + group.publish((None, None)) + torch.cuda._sleep(2_000_000) + with torch.cuda.stream(alias): + group.publish((8, 8)) + observed_first = [buf.gpu.clone() for buf in buffers] + with pytest.raises(PublicationError, match="republish_reason"): + buffers[0].copy_to_gpu() + owner.finish() + # begin must wait for old source reads and refresh identity on a new stream. + with torch.cuda.stream(external): + owner.begin() + group.publish((8, 8)) + owner.finish() + with torch.cuda.stream(second): + owner.begin() + for i, buf in enumerate(buffers): + buf.cpu.fill_(21 + i) + with ( + torch.cuda.stream(first), + pytest.raises(PublicationError, match="compute stream"), + ): + group.publish((8, 8)) + group.publish((8, 8)) + observed_second = [buf.gpu.clone() for buf in buffers] + owner.finish() + owner.completion.synchronize() + assert [buf.cpu().tolist() for buf in observed_first] == [[11] * 8, [12] * 8] + assert [buf.cpu().tolist() for buf in observed_second] == [[21] * 8, [22] * 8] + + +@gpu +@pytest.mark.parametrize("fast", [False, True], ids=["public", "tuple"]) +@pytest.mark.parametrize("backend", ["direct", "packed"]) +def test_capture_and_other_thread_reject_without_consuming_publication(fast, backend): + owner, buffers, group = stream_setup(fast, backend) + owner.begin() + group.publish((8, 8)) + owner.finish() + owner.begin() + for i, buf in enumerate(buffers): + buf.cpu.fill_(31 + i) + with ( + ThreadPoolExecutor(max_workers=1) as worker, + pytest.raises(PublicationError, match="owner thread"), + ): + worker.submit(group.publish, (8, 8)).result() + output = torch.empty_like(buffers[0].gpu) + graph = torch.cuda.CUDAGraph() + torch.cuda.synchronize() + with torch.cuda.graph(graph): + output.copy_(buffers[0].gpu) + with pytest.raises(PublicationError, match="actual graph capture"): + group.publish((8, 8)) + group.publish((8, 8)) + owner.finish() + graph.replay() + torch.cuda.synchronize() + assert output.cpu().tolist() == [31] * 8 + + +@gpu +@pytest.mark.parametrize("fast", [False, True], ids=["public", "tuple"]) +def test_other_device_rejects_before_publication(fast): + if torch.cuda.device_count() < 2: + pytest.skip("requires two GPU devices") + with torch.cuda.device(0): + owner, buffers, group = stream_setup(fast, "direct") + owner.begin() + buffers[0].cpu.fill_(41) + buffers[1].cpu.fill_(42) + with ( + torch.cuda.device(1), + pytest.raises(PublicationError, match="wrong device"), + ): + group.publish((8, 8)) + group.publish((8, 8)) + owner.finish() + owner.completion.synchronize() + assert [buf.gpu.cpu().tolist() for buf in buffers] == [[41] * 8, [42] * 8] diff --git a/tests/test_h2d_runner_publication.py b/tests/test_h2d_runner_publication.py new file mode 100644 index 0000000000..3cde73eb20 --- /dev/null +++ b/tests/test_h2d_runner_publication.py @@ -0,0 +1,587 @@ +# SPDX-License-Identifier: MIT +"""Real runner producers, variable counts and asynchronous staging reuse.""" + +import os +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from atom.utils import CpuGpuBuffer +from atom.utils.h2d import PublicationError + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1", reason="set RUN_H2D_GPU_TESTS=1" +) + + +def runner_with_buffers(monkeypatch, transport, pp_size=1, speculative=False): + from atom.model_engine.model_runner import ModelRunner, tokenIDProcessor + from atom.model_ops.attentions.qwen4_exp_attn import Qwen4ExpMetadataBuilder + + monkeypatch.setenv("ATOM_H2D_BACKEND", transport) + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace( + pipeline_parallel_size=pp_size, + max_num_batched_tokens=64, + max_num_seqs=8, + hf_config=SimpleNamespace(hidden_size=8), + torch_dtype=torch.float32, + ) + runner.model = SimpleNamespace() + runner.use_mrope = True + runner.enforce_eager = True + runner.tokenID_processor = tokenIDProcessor(runner, 64) + if speculative: + from atom.spec_decode.drafter import Drafter + + class MetadataDrafter(Drafter): + def _resolve_mtp_k(self): + return 3 + + def propose(self, *args): + raise NotImplementedError + + runner.drafter = MetadataDrafter.__new__(MetadataDrafter) + runner.drafter.runner = runner + runner.drafter.mtp_k = 3 + runner.drafter.metadata_buffers = Drafter._allocate_metadata_buffers( + 8, 3, runner.device + ) + runner.arange_np = np.arange(64, dtype=np.int32) + runner.allocate_forward_vars() + runner.forward_vars["cu_seqlens_q"] = CpuGpuBuffer( + 9, dtype=torch.int32, device=runner.device, publication_group="early" + ) + builder = Qwen4ExpMetadataBuilder.__new__(Qwen4ExpMetadataBuilder) + builder.model_runner = runner + runner.attn_metadata_builder = builder + runner._init_forward_vars_ring() + runner._init_h2d_publication() + return runner + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_sampling_optional_members_and_scalar_filters(monkeypatch, transport, pp_size): + from atom.model_ops.sampler import SAMPLER_EPS + + runner = runner_with_buffers(monkeypatch, transport, pp_size) + results = [] + cases = ( + ([0, 0, 0], [-1, -1, -1], [1, 1, 1]), + ([0.2, 0.4, 0.6], [4, 4, 4], [0.9, 0.9, 0.9]), + ([1, 0], [8, 2], [0.5, 0.8]), + ([0], [-1], [1]), + ([0.7], [4], [0.9]), + ([0.5, 1], [4, 4], [0.7, 0.8]), + ([0.5, 1], [4, 7], [0.8, 0.8]), + ) * 2 + for temps, ks, ps in cases: + runner._advance_forward_vars() + runner._gate_staging_reuse() + for name in ("temperatures", "top_ks", "top_ps"): + runner.forward_vars[name].gpu.fill_(-99) + batch = SimpleNamespace( + total_seqs_num=len(temps), + temperatures=np.array(temps), + top_ks=np.array(ks), + top_ps=np.array(ps), + needs_independent_noise=np.array([True] * len(temps)), + ) + torch.cuda._sleep(2_000_000) + t, k, p, greedy, noise = runner.prepare_sample(batch) + results.append( + ( + t.clone(), + k.clone() if isinstance(k, torch.Tensor) else k, + p.clone() if isinstance(p, torch.Tensor) else p, + greedy, + noise, + ) + ) + counts = runner.h2d_groups["sampling"].counts + for name in ("top_ks", "top_ps"): + count = counts[runner.h2d_groups["sampling"].indices[name]] + buf = runner.forward_vars[name] + # The omitted or unselected tail is never overwritten. + results[-1] += (buf.gpu[count or 0 :].clone(),) + with pytest.raises(PublicationError, match="republish_reason"): + runner.forward_vars["temperatures"].copy_to_gpu(len(temps)) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for (temps, ks, ps), (t, k, p, greedy, noise, ktail, ptail) in zip(cases, results): + np.testing.assert_array_equal( + t.cpu().numpy().view(np.uint8), + np.maximum(temps, SAMPLER_EPS).astype(np.float32).view(np.uint8), + ) + assert greedy == (np.array(temps) == 0).all() and noise + if all(value == -1 for value in ks): + assert k is None + elif len(set(ks)) == 1: + assert type(k) is int and k == ks[0] + else: + assert k.cpu().tolist() == ks + if all(value == 1 for value in ps): + assert p is None + elif len(set(ps)) == 1: + assert type(p) is float and p == float(np.float32(ps[0])) + else: + np.testing.assert_array_equal( + p.cpu().numpy().view(np.uint8), + np.asarray(ps, dtype=np.float32).view(np.uint8), + ) + assert torch.all(ktail.cpu() == -99) and torch.all(ptail.cpu() == -99) + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +def test_input_ids_deferred_and_new_requests_share_one_publication( + monkeypatch, transport +): + runner = runner_with_buffers(monkeypatch, transport) + processor = runner.tokenID_processor + runner._gate_staging_reuse() + processor.prev_batch = SimpleNamespace(req_ids=[10, 20], is_dummy_run=False) + processor.prev_token_ids = torch.tensor([71, 82], dtype=torch.int32, device="cuda") + batch = SimpleNamespace( + scheduled_tokens=np.array([-1, 93, -1], dtype=np.int32), + total_tokens_num=3, + total_tokens_num_prefill=0, + total_tokens_num_decode=3, + total_seqs_num_prefill=0, + total_seqs_num_decode=3, + total_seqs_num=3, + req_ids=[20, 30, 10], + is_dummy_run=False, + num_scheduled_tokens=np.ones(3, dtype=np.int32), + num_rejected=np.zeros(3, dtype=np.int32), + num_bonus=np.zeros(3, dtype=np.int32), + produces_output=lambda: True, + ) + runner.attn_metadata_builder.publish_cu_seqlens_q( + batch, SimpleNamespace(running_bs=3) + ) + torch.cuda._sleep(2_000_000) + ids = processor.prepare_input_ids(batch, 1).clone() + with pytest.raises(PublicationError, match="republish_reason"): + processor.input_ids.copy_to_gpu(3) + runner._mark_staging_h2d_enqueued() + runner.h2d_owner.completion.synchronize() + assert ids.cpu().tolist() == [82, 93, 71] + + +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_prefill_and_first_decode_ids_rotate_with_the_runner(monkeypatch, pp_size): + runner = runner_with_buffers(monkeypatch, "direct", pp_size) + outputs = [] + for i in range(6): + runner._advance_forward_vars() + runner._gate_staging_reuse() + processor = runner.tokenID_processor + assert processor.decode_src is runner.forward_vars["decode_src"] + prefill = i % 2 == 0 + batch = SimpleNamespace( + scheduled_tokens=np.array([i + 1, i + 2], dtype=np.int32), + total_tokens_num=2, + total_tokens_num_prefill=2 if prefill else 0, + total_tokens_num_decode=0 if prefill else 2, + total_seqs_num_prefill=1 if prefill else 0, + produces_output=lambda: False, + ) + torch.cuda._sleep(2_000_000) + outputs.append(processor.prepare_input_ids(batch, 1).clone()) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + assert [out.cpu().tolist() for out in outputs] == [[i + 1, i + 2] for i in range(6)] + + +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_mrope_padding_is_final_before_its_only_publication(monkeypatch, pp_size): + runner = runner_with_buffers(monkeypatch, "packed", pp_size) + outputs = [] + for i, bs in enumerate((4, 3, 1, 4, 2, 1)): + runner._advance_forward_vars() + runner._gate_staging_reuse() + buf = runner.forward_vars["mrope_positions"] + assert buf._publication.unit == "elements" + buf.gpu.fill_(-999) + batch = SimpleNamespace( + total_tokens_num_decode=bs * 2, + req_ids=range(bs), + mrope_position_deltas={j: j * 10 for j in range(bs)}, + ) + ends = np.arange(bs) * 100 + i * 2 + 2 + torch.cuda._sleep(2_000_000) + positions = runner.attn_metadata_builder._build_mrope_decode_positions( + batch, ends, 2, running_tokens=8 + ) + assert positions.stride(0) == 8 + outputs.append( + ( + positions.clone(), + runner._mrope_positions_view(8).clone(), + buf.gpu.reshape(-1)[24:].clone(), + bs, + i, + ) + ) + with pytest.raises(PublicationError, match="republish_reason"): + runner.attn_metadata_builder._copy_mrope_to_gpu(8) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for positions, padded, tail, bs, i in outputs: + expected = [j * 110 + i * 2 + k for j in range(bs) for k in range(2)] + assert positions.cpu().tolist() == [expected] * 3 + assert torch.all(padded.cpu()[:, bs * 2 :] == 0) + assert torch.all(tail.cpu() == -999) + + +@pytest.mark.parametrize( + "transport,coalesce", + [("direct", False), ("packed", False), ("packed", True)], +) +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_speculative_indices_publish_before_index_select( + monkeypatch, transport, coalesce, pp_size +): + from atom.spec_decode.drafter import Drafter + + runner = runner_with_buffers(monkeypatch, transport, pp_size, speculative=True) + scratch = runner.forward_vars["verification_draft_token_ids"] + saved = [] + for step, lengths in enumerate(([3, 1, 2], [1, 1], [4] * 8, [1, 3, 1]) * 2): + runner._advance_forward_vars() + runner._gate_staging_reuse() + # Device writes and sampler reads are ordered on the forward stream, + # including across PP slots; no pinned host scratch is involved. + assert runner.forward_vars["verification_draft_token_ids"] is scratch + scratch.fill_(-123) + lengths = np.array(lengths, dtype=np.int32) + ends = np.cumsum(lengths) + ids = torch.arange(int(ends[-1]), device="cuda", dtype=torch.int32) + step * 100 + torch.cuda._sleep(2_000_000) + prepared = None + if coalesce: + group = runner.h2d_groups["token_inputs"] + prepared = runner.drafter.prepare_spec_decode_indices(lengths, ends, group) + group.publish(group.counts) + metadata = runner.drafter.calc_spec_decode_metadata( + lengths, ends, ids, prepared_indices=prepared + ) + if coalesce: + assert runner.h2d_groups["spec_decode"]._backend is None + assert ( + metadata.draft_token_ids.untyped_storage().data_ptr() == scratch.data_ptr() + ) + saved_tail = scratch[metadata.draft_token_ids.numel() :].clone() + anchors = Drafter.anchors_to_gpu(runner.drafter, [-1] * len(lengths)) + saved.append( + ( + metadata.draft_token_ids.clone(), + metadata.target_logits_indices.clone(), + metadata.bonus_logits_indices.clone(), + metadata.cu_num_draft_tokens.clone(), + anchors.clone(), + lengths, + step, + saved_tail, + ) + ) + # Even a zero-draft step explicitly published its empty index prefix. + with pytest.raises(PublicationError, match="republish_reason"): + runner.forward_vars["target_logits_indices"].copy_to_gpu(0) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for drafts, targets, bonus, cu, anchors, lengths, step, tail in saved: + offsets = np.cumsum(lengths) - lengths + expected_targets = [ + int(start + j) + for start, length in zip(offsets, lengths) + for j in range(int(length) - 1) + ] + assert targets.cpu().tolist() == expected_targets + assert drafts.cpu().tolist() == [step * 100 + j + 1 for j in expected_targets] + assert bonus.cpu().tolist() == (np.cumsum(lengths) - 1).tolist() + assert cu.cpu().tolist() == np.cumsum(lengths - 1).tolist() + assert anchors.cpu().tolist() == [-1] * len(lengths) + assert torch.all(tail.cpu() == -123) + + +@pytest.mark.parametrize("phase", ["prefill", "first_decode", "deferred"]) +@pytest.mark.parametrize("transport", ["direct", "packed"]) +def test_prepare_model_publishes_sampling_and_query_prefix_before_token_consumer( + monkeypatch, phase, transport +): + from atom.model_engine import model_runner + + runner = runner_with_buffers(monkeypatch, transport) + runner.config.parallel_config = SimpleNamespace(data_parallel_size=1) + runner.config.enable_tbo = False + runner.capture_sizes_np = np.array([1, 2, 4, 8]) + runner._dspark_apply_q_bucket = lambda batch: None + runner._piecewise_cg_active = lambda: False + runner._local_tbo_eligibility = lambda batch: False + runner.prepare_inputs = lambda *args, **kwargs: None + mode = SimpleNamespace(sync=None, max_seqlen_q=1, running_bs=4) + monkeypatch.setattr(model_runner.ForwardMode, "decide", lambda **kwargs: mode) + processor = runner.tokenID_processor + prefill = phase == "prefill" + if phase == "deferred": + processor.prev_batch = SimpleNamespace(req_ids=[10, 20], is_dummy_run=False) + processor.prev_token_ids = torch.tensor( + [71, 82], dtype=torch.int32, device="cuda" + ) + batch = SimpleNamespace( + scheduled_tokens=np.array([51, 93, 61], dtype=np.int32), + total_tokens_num=3, + total_tokens_num_prefill=3 if prefill else 0, + total_tokens_num_decode=0 if prefill else 3, + total_seqs_num_prefill=3 if prefill else 0, + total_seqs_num_decode=0 if prefill else 3, + total_seqs_num=3, + num_spec_step=0, + req_ids=[20, 30, 10], + is_dummy_run=False, + num_scheduled_tokens=np.ones(3, dtype=np.int32), + num_rejected=np.zeros(3, dtype=np.int32), + num_bonus=np.zeros(3, dtype=np.int32), + temperatures=np.array([0.5, 0.7, 0.9], dtype=np.float32), + top_ks=np.array([4, 7, 9], dtype=np.int32), + top_ps=np.array([0.8, 0.9, 0.7], dtype=np.float32), + produces_output=lambda: True, + ) + for step in range(3): + runner._gate_staging_reuse() + torch.cuda._sleep(2_000_000) + ids, temperatures, ks, ps, _, _ = runner.prepare_model(batch) + observed = [x.clone() for x in (ids, temperatures, ks, ps)] + cu = runner.forward_vars["cu_seqlens_q"].gpu[:5].clone() + runner._mark_staging_h2d_enqueued() + runner.h2d_owner.completion.synchronize() + assert observed[0].cpu().tolist() == ( + [82, 93, 71] if phase == "deferred" else [51, 93, 61] + ) + for value, expected in zip( + observed[1:], (batch.temperatures, batch.top_ks, batch.top_ps) + ): + np.testing.assert_array_equal(value.cpu().numpy(), expected) + assert cu.cpu().tolist() == [0, 1, 2, 3, 3] + if transport == "packed": + group = runner.h2d_groups["token_inputs"] + assert group._backend.kernel is not None + assert all( + b._epoch == runner.h2d_owner.epoch + for b in group.members + if b.name != "decode_src" or phase == "deferred" + ) + + +def test_packed_spec_indices_follow_decode_prefill_and_dummy_transitions(monkeypatch): + from atom.model_engine import model_runner + + runner = runner_with_buffers(monkeypatch, "packed", speculative=True) + runner.config.parallel_config = SimpleNamespace(data_parallel_size=1) + runner.config.enable_tbo = False + runner.capture_sizes_np = np.array([1, 2, 4, 8]) + runner._dspark_apply_q_bucket = lambda batch: None + runner._piecewise_cg_active = lambda: False + runner._local_tbo_eligibility = lambda batch: False + mode = SimpleNamespace(sync=None, max_seqlen_q=4, running_bs=4) + monkeypatch.setattr(model_runner.ForwardMode, "decide", lambda **kwargs: mode) + monkeypatch.setattr(model_runner, "get_forward_context", lambda: None) + saved = [] + + def consume(batch, ids, forward_mode, *, spec_decode_indices): + if batch.total_tokens_num_prefill or batch.is_dummy_run: + assert spec_decode_indices is None + group = runner.h2d_groups["token_inputs"] + for binding in runner.h2d_groups["spec_decode"].members: + assert group.counts[group.indices[binding.name]] is None + assert binding._epoch != runner.h2d_owner.epoch + return + _, lens, cu = runner.attn_metadata_builder.decode_spans(batch) + metadata = runner.drafter.calc_spec_decode_metadata( + lens, cu[1:], ids, prepared_indices=spec_decode_indices + ) + saved.append(metadata.draft_token_ids.clone()) + assert runner.h2d_groups["spec_decode"]._backend is None + + runner.prepare_inputs = consume + for phase in ["decode", "prefill", "dummy", "decode"]: + prefill, dummy = phase == "prefill", phase == "dummy" + batch = SimpleNamespace( + scheduled_tokens=np.arange(7, dtype=np.int32) + 10, + total_tokens_num=7, + total_tokens_num_prefill=7 if prefill else 0, + total_tokens_num_decode=0 if prefill else 7, + total_seqs_num_prefill=3 if prefill else 0, + total_seqs_num_decode=0 if prefill else 3, + total_seqs_num=3, + num_spec_step=3, + req_ids=[20, 30, 10], + is_dummy_run=dummy, + num_scheduled_tokens=np.array([4, 1, 2], dtype=np.int32), + temperatures=np.zeros(3, dtype=np.float32), + top_ks=np.full(3, -1, dtype=np.int32), + top_ps=np.ones(3, dtype=np.float32), + produces_output=lambda: True, + next_token_ids=None, + ) + runner._gate_staging_reuse() + torch.cuda._sleep(2_000_000) + runner.prepare_model(batch) + runner._mark_staging_h2d_enqueued() + runner.h2d_owner.completion.synchronize() + assert len(saved) == 2 + assert all(value.cpu().tolist() == [11, 12, 13, 16] for value in saved) + + +@pytest.mark.parametrize("builder_kind", ["common", "qwen4"]) +def test_mrope_producer_and_padded_model_view_share_axis_stride( + monkeypatch, builder_kind +): + from atom.model_ops.attentions.aiter_attention import AiterAttentionMetadataBuilder + + runner = runner_with_buffers(monkeypatch, "direct") + builder = runner.attn_metadata_builder + if builder_kind == "common": + builder.__class__ = AiterAttentionMetadataBuilder + runner._gate_staging_reuse() + batch = SimpleNamespace( + total_tokens_num_decode=4, + req_ids=[10, 20], + mrope_position_deltas={10: 100, 20: 200}, + ) + builder._build_mrope_decode_positions( + batch, np.array([12, 22]), 2, running_tokens=8 + ) + _, positions = runner._padded_decode_inputs( + SimpleNamespace( + is_prefill=False, + running_tokens_are_unified=True, + scheduled_tokens=4, + running_tokens=8, + ) + ) + observed = positions.clone() + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + assert observed.tolist() == [[110, 111, 220, 221, 0, 0, 0, 0]] * 3 + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +@pytest.mark.parametrize("sealed", [False, True], ids=["active", "sealed"]) +@pytest.mark.parametrize( + "producer", + ["sampling", "prefill_ids", "deferred_ids", "query_prefix", "spec_indices"], +) +def test_producer_reentry_preserves_sources_and_first_consumer( + monkeypatch, transport, sealed, producer +): + runner = runner_with_buffers(monkeypatch, transport, speculative=True) + processor = runner.tokenID_processor + combined = runner.h2d_groups.get("token_inputs") + source_name = { + "sampling": "sampling", + "prefill_ids": "input_ids", + "deferred_ids": "input_ids", + "query_prefix": "early", + "spec_indices": "spec_decode", + }[producer] + sources = runner.h2d_groups[source_name] + if producer == "deferred_ids": + processor.prev_batch = SimpleNamespace(req_ids=[10, 20], is_dummy_run=False) + processor.prev_token_ids = torch.tensor( + [71, 82], dtype=torch.int32, device="cuda" + ) + + def produce(changed): + if producer == "sampling": + batch = SimpleNamespace( + total_seqs_num=3, + temperatures=np.array([0.2, 0.4, 0.6]) + changed, + top_ks=np.array([4, 5, 6]) + changed, + top_ps=np.array([0.7, 0.8, 0.9]) - changed * 0.1, + ) + runner.prepare_sample(batch, publication_group=combined) + if combined is not None: + combined.publish(combined.counts) + elif producer in ("prefill_ids", "deferred_ids"): + prefill = producer == "prefill_ids" + batch = SimpleNamespace( + scheduled_tokens=np.array([11, 12, 13], dtype=np.int32) + changed * 10, + total_tokens_num=3, + total_tokens_num_prefill=3 if prefill else 0, + total_tokens_num_decode=0 if prefill else 3, + total_seqs_num_prefill=3 if prefill else 0, + total_seqs_num_decode=0 if prefill else 3, + total_seqs_num=3, + req_ids=[20, 30, 10] if not changed else [10, 20, 30], + is_dummy_run=False, + num_scheduled_tokens=np.ones(3, dtype=np.int32), + num_rejected=np.zeros(3, dtype=np.int32), + num_bonus=np.zeros(3, dtype=np.int32), + produces_output=lambda: True, + ) + processor.prepare_input_ids(batch, 1, publication_group=combined) + elif producer == "query_prefix": + lengths = np.array([2, 3], dtype=np.int32) + changed + batch = SimpleNamespace( + total_seqs_num=2, + num_scheduled_tokens=lengths, + total_tokens_num=int(lengths.sum()), + ) + mode = SimpleNamespace(running_bs=3) + if combined is None: + runner.attn_metadata_builder.publish_cu_seqlens_q(batch, mode) + else: + count = runner.attn_metadata_builder.prepare_cu_seqlens_q(batch, mode) + combined.set_count(runner.forward_vars["cu_seqlens_q"], count) + combined.publish(combined.counts) + else: + lengths = np.array([2, 3], dtype=np.int32) + changed + group = sources if combined is None else combined + runner.drafter.prepare_spec_decode_indices( + lengths, np.cumsum(lengths), group + ) + group.publish(group.counts) + + # Deferred token assembly consumes an already-published unit-stride prefix. + # This immutable input is not among the source group being tested. + cu = runner.forward_vars["cu_seqlens_q"] + if producer == "deferred_ids": + cu.cpu[:4] = torch.arange(4, dtype=torch.int32, device="cpu") + cu.gpu[:4].copy_(cu.cpu[:4]) + runner._gate_staging_reuse() + produce(0) # Warm packing/token assembly before the delayed epoch. + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + expected = [member.destination.clone() for member in sources.members] + observed = [torch.empty_like(value) for value in expected] + + runner._gate_staging_reuse() + torch.cuda._sleep(20_000_000) + produce(0) + host_before = [member.source.clone() for member in sources.members] + for output, member in zip(observed, sources.members): + output.copy_(member.destination) + if sealed: + runner._mark_staging_h2d_enqueued() + try: + with pytest.raises(PublicationError, match="republish_reason|owner is sealed"): + produce(1) + finally: + if not sealed: + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + for member, before, output, reference in zip( + sources.members, host_before, observed, expected + ): + assert torch.equal(member.source, before), member.name + assert torch.equal(output, reference), member.name diff --git a/tests/test_h2d_v4_indexer_publication.py b/tests/test_h2d_v4_indexer_publication.py new file mode 100644 index 0000000000..b316a46d17 --- /dev/null +++ b/tests/test_h2d_v4_indexer_publication.py @@ -0,0 +1,634 @@ +# SPDX-License-Identifier: MIT +"""V4 indexer/PCP producers use the forward slot through late TBO prepare.""" + +import os +import sys +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from atom.utils import CpuGpuBuffer +from atom.utils.h2d import PublicationError + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1", reason="set RUN_H2D_GPU_TESTS=1" +) + + +def make_runner(monkeypatch, *, pp_size=1, opus=False, base=None): + import atom.config + from atom.model_engine.model_runner import ModelRunner + from atom.model_ops.attentions.deepseek_v4_attn import ( + DeepseekV4AttentionMetadataBuilder, + ) + from atom.model_ops.attentions.pool_layout.v4_pool_fields import FP4_GFX1250_NATURAL + + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace( + pipeline_parallel_size=pp_size, enable_tbo=True, enable_tbo_decode=True + ) + monkeypatch.setattr(atom.config, "_current_atom_config", runner.config) + runner.enforce_eager = True + runner.tokenID_processor = SimpleNamespace() + runner.pool_plan = SimpleNamespace(entries={}) + runner.forward_vars = {} + for name, capacity, dtype, group in ( + ("input_ids", 32, torch.int32, None), + ("decode_src", 32, torch.int32, None), + ("positions", 32, torch.int64, "positions"), + ("context_lens", 4, torch.int32, "context"), + ("cu_seqlens_q", 5, torch.int32, "early"), + ("cu_seqlens_k", 5, torch.int32, "cu_k"), + ("batch_id_per_q_token", 32, torch.int32, None), + ): + runner.forward_vars[name] = CpuGpuBuffer( + capacity, dtype=dtype, device=runner.device, publication_group=group + ) + runner.forward_vars["block_tables"] = CpuGpuBuffer( + 4, 4, dtype=torch.int32, device=runner.device, publication_group="blocks" + ) + cls = base or DeepseekV4AttentionMetadataBuilder + builder = cls.__new__(cls) + builder.model_runner = runner + builder.device = runner.device + builder.max_num_batched_tokens = 32 + builder.max_bs = 4 + builder.max_decode_tokens = 16 + builder.window_size = 4 + builder.index_topk = 4 + builder.max_committed_hca = 4 + builder.block_table_cols = 4 + builder.max_spec_steps = 2 + builder._unique_compress_ratios_overlap = [(4, True)] + builder._indexer_fp4 = False + builder.indexer_layout = FP4_GFX1250_NATURAL if opus else "fp8" + builder.pool_geometry = SimpleNamespace(slot_positions=4) + builder._alloc_v4_metadata_buffers() + runner.attn_metadata_builder = builder + monkeypatch.setenv("ATOM_H2D_BACKEND", "direct") + runner._init_forward_vars_ring() + runner._init_h2d_publication() + return runner, builder + + +def make_metadata(builder, lengths=(3, 5, 3), starts=(8, 16, 0), *, stage_map=True): + from atom.model_ops.attentions.deepseek_v4_attn import AttentionMetaData_DSV4 + from atom.utils.forward_context import AttentionMetaData, AttnState + from atom.utils.tbo.ubatch_splitting import attach_tbo_cpu_lens + + lengths = np.asarray(lengths, dtype=np.int32) + starts = np.asarray(starts, dtype=np.int32) + cu = np.r_[np.int32(0), np.cumsum(lengths, dtype=np.int32)] + positions = np.concatenate( + [ + np.arange(start, start + n, dtype=np.int64) + for start, n in zip(starts, lengths) + ] + ) + pos_gpu = builder._stage("positions", positions) + cu_gpu = builder._stage("cu_seqlens_q", cu) + ctx_gpu = builder._stage("context_lens", starts + lengths) + tables = builder._stage("block_tables", np.ones((len(lengths), 4), dtype=np.int32)) + slots = torch.arange(len(lengths), device="cuda", dtype=torch.int32) + md = AttentionMetaData( + cu_seqlens_q=cu_gpu, + cu_seqlens_k=cu_gpu, + max_seqlen_q=int(lengths.max()), + max_seqlen_k=int((starts + lengths).max()), + context_lens=ctx_gpu, + block_tables=tables, + state=AttnState.PREFILL_NATIVE, + has_cached=False, + ) + # Serving constructs the base metadata before promoting it to DSV4. + md.__class__ = AttentionMetaData_DSV4 + md.state_slot_out = slots + md.state_slot_in = slots + md.state_slot_out_cpu = np.arange(len(lengths), dtype=np.int32) + if stage_map: + builder._attach_v4_per_fwd_meta( + md, + np.repeat(np.arange(len(lengths), dtype=np.int32), lengths), + md.state_slot_out_cpu, + len(lengths), + len(positions), + ) + md.indexer_meta = {} + for name, value in ( + ("cu_seqlens_q", cu), + ("cu_seqlens_k", cu), + ("context_lens", starts + lengths), + ): + attach_tbo_cpu_lens(md, name, value.copy()) + return md, pos_gpu, cu, lengths + + +def set_pcp(monkeypatch, size, rank): + from atom.distributed.pcp_utils import pcp_round_robin_query_indices + from atom.model_ops.attentions import deepseek_v4_attn as module + + monkeypatch.setattr(module, "pcp_is_enabled", lambda: True) + monkeypatch.setattr(module, "get_pcp_world_size", lambda: size) + monkeypatch.setattr( + module, + "pcp_round_robin_query_indices", + lambda total, width: pcp_round_robin_query_indices(total, width, rank), + ) + + +@pytest.mark.parametrize("pp_size", [1, 2]) +@pytest.mark.parametrize("rank", [0, 3]) +def test_pcp_ragged_queries_publish_indexer_once(monkeypatch, pp_size, rank): + runner, builder = make_runner(monkeypatch, pp_size=pp_size) + set_pcp(monkeypatch, 4, rank) + saved = [] + for step in range(5): + runner._advance_forward_vars() + runner._gate_staging_reuse() + md, positions, cu, lengths = make_metadata( + builder, starts=(8 + 4 * step, 16, 0) + ) + md.kv_indptr_extend = torch.arange(12, dtype=torch.int32, device="cuda") + md.kv_indices_extend = torch.arange(11, dtype=torch.int32, device="cuda") + buf = runner.forward_vars["v4_indexer_cu_committed"] + buf.gpu.fill_(77) + torch.cuda._sleep(5_000_000) + local = builder._apply_pcp_reindex(md, positions, 3, 11, cu) + expected_cu = np.r_[0, np.cumsum(md.n_committed_csa_per_seq_cpu)].astype( + np.int32 + ) + expected_cu[-1] = max(expected_cu[-1], 1) + saved.append( + ( + md.indexer_meta["cu_committed_gpu"].clone(), + expected_cu, + md.batch_id_per_q_token.clone(), + md.kv_indices_extend.clone(), + local.clone(), + ) + ) + before = buf.cpu.clone() + with pytest.raises(PublicationError, match="republish_reason"): + builder._build_v4_indexer_meta( + attn_metadata=md, + positions_gpu=local, + scheduled_bs=3, + total_tokens=3, + device=runner.device, + ) + assert torch.equal(buf.cpu, before) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + ids = np.r_[np.repeat(np.arange(3), lengths), -1][rank::4] + for actual, expected, bids, indices, local in saved: + np.testing.assert_array_equal(actual.cpu().numpy(), expected) + np.testing.assert_array_equal(bids.cpu().numpy(), ids) + assert indices.cpu().tolist() == list(range(rank, 11, 4)) + assert len(local) == 3 + assert torch.all(buf.gpu[4:] == 77) + + +@pytest.mark.parametrize("balanced", [False, True]) +@pytest.mark.parametrize("dummy", [False, True]) +def test_prepare_prefill_defers_indexer_only_for_real_pcp(monkeypatch, balanced, dummy): + from atom.model_ops.attentions.backends import CommonAttentionBuilder + + runner, builder = make_runner(monkeypatch) + set_pcp(monkeypatch, 4, 0) + runner._pcp_tbo_balanced_active = balanced + runner._gate_staging_reuse() + md, pos, _, _ = make_metadata(builder, stage_map=False) + monkeypatch.setattr( + CommonAttentionBuilder, "prepare_prefill", lambda *args: (md, pos) + ) + batch = SimpleNamespace( + total_seqs_num_prefill=3, + total_tokens_num_prefill=11, + is_dummy_run=dummy, + state_slots_committed=[], + state_fork_srcs=None, + ) + actual, _ = builder.prepare_prefill(batch, 3) + binding = runner.forward_vars["v4_indexer_cu_committed"]._publication + assert (binding._epoch == runner.h2d_owner.epoch) == (dummy or not balanced) + assert len(actual.batch_id_per_q_token) == (11 if dummy or balanced else 3) + runner._mark_staging_h2d_enqueued() + + +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_token_split_tbo_resumes_epoch_and_preserves_both_prefixes( + monkeypatch, pp_size +): + from atom.utils.tbo.ubatch_splitting import UBatchSlice + + runner, builder = make_runner(monkeypatch, pp_size=pp_size) + saved = [] + for step in range(5): + runner._advance_forward_vars() + runner._gate_staging_reuse() + md, _, _, _ = make_metadata(builder, (5, 4), (8 + 4 * step, 20)) + runner._mark_staging_h2d_enqueued() + epoch = runner.h2d_owner.epoch + ubatches = [] + for index, sl in enumerate( + ( + UBatchSlice(slice(0, 1), slice(0, 4)), + UBatchSlice(slice(0, 2), slice(4, 9)), + ) + ): + torch.cuda._sleep(5_000_000) + ub = builder.build_ubatch_prefill_metadata( + md, sl, sl.request_slice.stop, index + ) + assert runner.h2d_owner.epoch == epoch + assert runner.h2d_owner._state == "sealed" + buf = runner.forward_vars[f"ub{index}_cu_seqlens_q"] + assert ub.cu_seqlens_q.data_ptr() == buf.gpu.data_ptr() + ubatches.append(ub) + # Read after BOTH preparations. A snapshot taken before the second + # producer could conceal a shared-storage overwrite. + for index, ub in enumerate(ubatches): + assert ( + ub.batch_id_per_q_token.data_ptr() + == runner.forward_vars[f"ub{index}_batch_id_per_q_token"].gpu.data_ptr() + ) + assert ( + ub.indexer_meta["cu_committed_gpu"].data_ptr() + == runner.forward_vars[ + f"ub{index}_v4_indexer_cu_committed" + ].gpu.data_ptr() + ) + saved.append( + ( + index, + ub.cu_seqlens_q.clone(), + ub.indexer_meta["cu_committed_gpu"].clone(), + step, + ) + ) + before = runner.forward_vars["ub0_context_lens"].cpu.clone() + with pytest.raises(PublicationError, match="republish_reason"): + builder.build_ubatch_prefill_metadata( + md, UBatchSlice(slice(0, 1), slice(0, 4)), 1, 0 + ) + assert runner.h2d_owner._state == "sealed" + assert torch.equal(runner.forward_vars["ub0_context_lens"].cpu, before) + runner._record_forward_vars_event() + torch.cuda.synchronize() + for index, query, committed, step in saved: + assert query.cpu().tolist() == ([0, 4] if index == 0 else [0, 1, 5]) + assert committed.cpu().tolist() == ( + [0, 3 + step] if index == 0 else [0, 3 + step, 9 + step] + ) + + +@pytest.mark.parametrize("rank", [0, 3]) +def test_balanced_pcp_tbo_uses_separate_indexer_buffers(monkeypatch, rank): + runner, builder = make_runner(monkeypatch) + set_pcp(monkeypatch, 4, rank) + runner._pcp_tbo_balanced_active = True + runner._pcp_bal_groups = [ + SimpleNamespace(req_start=0, req_stop=2, tok_start=0, tok_end=8), + SimpleNamespace(req_start=2, req_stop=4, tok_start=8, tok_end=15), + ] + runner._gate_staging_reuse() + md, _, _, _ = make_metadata(builder, (3, 5, 3, 4), (8, 16, 28, 40)) + runner._mark_staging_h2d_enqueued() + ubatches = [] + for index in range(2): + torch.cuda._sleep(5_000_000) + ubatches.append(builder.build_ubatch_prefill_metadata(md, None, 2, index)) + runner.h2d_owner.completion.synchronize() + assert ubatches[0].indexer_meta["cu_committed_gpu"].cpu().tolist() == [0, 2, 7] + assert ubatches[1].indexer_meta["cu_committed_gpu"].cpu().tolist() == [0, 7, 18] + assert [ub.max_seqlen_q for ub in ubatches] == [5, 4] + for index, ub in enumerate(ubatches): + assert ( + ub.indexer_meta["cu_committed_gpu"].data_ptr() + == runner.forward_vars[f"ub{index}_v4_indexer_cu_committed"].gpu.data_ptr() + ) + assert runner.forward_vars["v4_indexer_cu_committed"]._publication._epoch == -1 + for index in range(2): + assert ( + runner.forward_vars[ + f"ub{index}_v4_indexer_cu_committed" + ]._publication._epoch + == runner.h2d_owner.epoch + ) + + +def mock_opus(monkeypatch): + name = "aiter.ops.opus.pa_mqa_logits_mxfp4" + + def plan(cu, ends, **kwargs): + return (cu.clone(), ends.clone(), kwargs["row_to_batch"].clone()) + + monkeypatch.setitem( + sys.modules, name, SimpleNamespace(pa_mqa_logits_mxfp4_plan=plan) + ) + + +@pytest.mark.parametrize("prefix", ["", "ub0_", "ub1_"]) +def test_opus_chunks_are_one_checked_upload_with_independent_ranges( + monkeypatch, prefix +): + from atom.model_ops.attentions import deepseek_v4_attn as module + + mock_opus(monkeypatch) + monkeypatch.setattr(module, "sparse_indexer_row_chunk", lambda *args: 2) + runner, builder = make_runner(monkeypatch, opus=True) + runner._gate_staging_reuse() + buf = runner.forward_vars[f"{prefix}v4_indexer_chunk_cu"] + buf.gpu.fill_(77) + cu = np.array([0, 0, 3, 3, 8], dtype=np.int32) + md = SimpleNamespace( + n_committed_csa_per_seq_cpu=np.array([0, 2, 0, 4], dtype=np.int32), + batch_id_per_q_token=torch.tensor( + [1, 1, 1, 3, 3, 3, 3, 3], device="cuda", dtype=torch.int32 + ), + ) + ends = torch.arange(8, dtype=torch.int32, device="cuda") + meta = {} + torch.cuda._sleep(5_000_000) + builder._build_fp4_opus_prefill_plans( + attn_metadata=md, + meta=meta, + total_tokens=8, + scheduled_bs=4, + visible_end_gpu=ends, + cu_seqlens_q_cpu=cu, + reuse_cu_seqlens_q=False, + plan_total_tokens=None, + buf_prefix_ubatch=prefix, + ) + assert buf._publication._epoch == runner.h2d_owner.epoch + runner._mark_staging_h2d_enqueued() + runner.h2d_owner.completion.synchronize() + entries = 0 + for start, end, plan in meta["fp4_opus_prefill_chunks"]: + expected = module._chunk_cu_seqlens(cu, start, end) + np.testing.assert_array_equal(plan[0].cpu().numpy(), expected) + entries += len(expected) + assert torch.all(buf.gpu[entries:] == 77) + + +def test_opus_full_prefix_reuse_does_not_publish_chunk_storage(monkeypatch): + from atom.model_ops.attentions import deepseek_v4_attn as module + + mock_opus(monkeypatch) + monkeypatch.setattr(module, "sparse_indexer_row_chunk", lambda *args: 32) + runner, builder = make_runner(monkeypatch, opus=True) + runner._gate_staging_reuse() + md, pos, cu, _ = make_metadata(builder) + meta = {} + builder._build_fp4_opus_prefill_plans( + attn_metadata=md, + meta=meta, + total_tokens=11, + scheduled_bs=3, + visible_end_gpu=pos.to(torch.int32), + cu_seqlens_q_cpu=cu, + reuse_cu_seqlens_q=True, + plan_total_tokens=None, + ) + assert runner.forward_vars["v4_indexer_chunk_cu"]._publication._epoch == -1 + assert torch.equal(meta["fp4_opus_prefill_chunks"][0][2][0], md.cu_seqlens_q) + runner._mark_staging_h2d_enqueued() + + +def test_opus_packed_source_reuse_and_capacity_check(monkeypatch): + from atom.model_ops.attentions import deepseek_v4_attn as module + + mock_opus(monkeypatch) + monkeypatch.setattr(module, "sparse_indexer_row_chunk", lambda *args: 1) + runner, builder = make_runner(monkeypatch, opus=True) + saved = [] + for step in range(5): + runner._gate_staging_reuse() + md, pos, cu, _ = make_metadata(builder, (step + 1, 4), (8, 20)) + buf = runner.forward_vars["v4_indexer_chunk_cu"] + before = buf.cpu.clone() + kwargs = { + "attn_metadata": md, + "total_tokens": len(pos), + "scheduled_bs": 2, + "visible_end_gpu": pos.to(torch.int32), + "cu_seqlens_q_cpu": cu, + "reuse_cu_seqlens_q": False, + } + with pytest.raises(ValueError, match="capacity"): + builder._build_fp4_opus_prefill_plans( + meta={}, plan_total_tokens=100, **kwargs + ) + assert torch.equal(buf.cpu, before) + assert buf._publication._epoch != runner.h2d_owner.epoch + torch.cuda._sleep(5_000_000) + meta = {} + builder._build_fp4_opus_prefill_plans( + meta=meta, plan_total_tokens=None, **kwargs + ) + saved.append(meta["fp4_opus_prefill_chunks"]) + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + for chunks in saved: + assert all(plan[0].cpu().tolist() == [0, 1] for _, _, plan in chunks) + + +def test_late_ubatch_failure_poisoning_and_wrong_stream_rejection(monkeypatch): + from atom.utils.tbo.ubatch_splitting import UBatchSlice + + runner, builder = make_runner(monkeypatch) + runner._gate_staging_reuse() + md, _, _, _ = make_metadata(builder, (5, 4), (8, 20)) + runner._mark_staging_h2d_enqueued() + sl = UBatchSlice(slice(0, 1), slice(0, 4)) + with ( + torch.cuda.stream(torch.cuda.Stream()), + pytest.raises(PublicationError, match="original owner stream"), + ): + builder.build_ubatch_prefill_metadata(md, sl, 1, 0) + assert runner.h2d_owner._state == "sealed" + + def fail(*args, **kwargs): + raise RuntimeError("indexer consumer failed") + + monkeypatch.setattr(builder, "_attach_v4_indexer_meta", fail) + with pytest.raises(RuntimeError, match="consumer failed"): + builder.build_ubatch_prefill_metadata(md, sl, 1, 0) + assert runner.h2d_owner._state == "failed" + before = runner.forward_vars["ub0_context_lens"].cpu.clone() + with pytest.raises(PublicationError, match="failed"): + builder.build_ubatch_prefill_metadata(md, sl, 1, 0) + assert torch.equal(runner.forward_vars["ub0_context_lens"].cpu, before) + runner.h2d_owner.drain() + + +@pytest.mark.parametrize("pp_size", [1, 2]) +@pytest.mark.parametrize("padded", [False, True]) +@pytest.mark.parametrize("kv_fp8", [False, True]) +def test_decode_ubatch_inputs_padding_and_graph_consumers( + monkeypatch, pp_size, padded, kv_fp8 +): + from atom.model_ops.attentions import deepseek_v4_attn as module + from atom.model_ops.attentions.pool_layout.v4_pool_geometry import ( + UnifiedPoolGeometry, + ) + from atom.utils.tbo.ubatch_wrapper import UBatchWrapper + + runner, builder = make_runner(monkeypatch, pp_size=pp_size) + builder.pool_geometry = UnifiedPoolGeometry([0, 4, 128], 8, 4, 6, 256) + builder.hca_rows_per_block = 2 + builder._kv_fp8 = kv_fp8 + runner.enforce_eager = not padded + if padded: + monkeypatch.setattr(module, "get_forward_context", lambda: None) + monkeypatch.setattr(UBatchWrapper, "_decode_ub_running_bs", lambda *args: 2) + names = [ + f"ub{i}_{name}" + for i in range(2) + for name in ( + "positions", + "context_lens", + "cu_seqlens_q", + "block_tables", + "v4_meta_state_slot_out", + "v4_meta_state_slot_in", + "batch_id_per_q_token", + ) + ] + if kv_fp8: + names.extend(f"ub{i}_v4_qo_indptr" for i in range(2)) + graphs, outputs = [], [] + for var in runner._fv_ring: + out = {name: torch.empty_like(var[name].gpu) for name in names} + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for name in names: + out[name].copy_(var[name].gpu) + graphs.append(graph) + outputs.append(out) + saved = [] + for step, lens in enumerate(((1, 3, 2), (2,), (), (3, 1, 2, 1), (1, 2, 3))): + runner._advance_forward_vars() + runner._gate_staging_reuse() + var = runner.forward_vars + for name in names: + var[name].gpu.fill_(77) + # Test-only storage poisoning also discards the published revision. + if hasattr(var[name], "_block_table"): + del var[name]._block_table + lengths = np.asarray(lens, dtype=np.int32) + starts = np.arange(len(lens), dtype=np.int32) * 8 + 8 + step + positions = ( + np.concatenate([np.arange(s, s + n) for s, n in zip(starts, lens)]) + if lens + else np.empty(0, dtype=np.int64) + ) + slots = np.arange(len(lens), dtype=np.int32) + var["block_tables"].np[: len(lens)] = slots[:, None] + 1 + kwargs = { + "scheduled_bs": len(lens), + "running_bs": 4 if padded else len(lens), + "max_seqlen_q": 3, + "context_lens_np": starts + lengths, + "state_slot_np": slots, + "state_slot_in_np": slots[::-1].copy(), + "positions_np": positions, + "extend_lens_np": lengths, + } + global_qo = np.minimum(np.arange(13, dtype=np.int32), len(positions)) + if kv_fp8: + builder._stage("v4_qo_indptr", global_qo) + torch.cuda._sleep(5_000_000) + builder._prepare_ubatch_decode(**kwargs) + if kv_fp8: + # The actual serving path prepared the global verify prefix first. + np.testing.assert_array_equal(var["v4_qo_indptr"].np[:13], global_qo) + saved.append((var["v4_qo_indptr"].gpu[:13].clone(), global_qo.tolist())) + graphs[runner._fv_idx].replay() + split = min(len(lens), 2) if padded else len(lens) // 2 + for index, (lo, hi) in enumerate(((0, split), (split, len(lens)))): + width = 2 if padded else hi - lo + local_lens = lengths[lo:hi] + tokens = int(local_lens.sum()) + local_pos = positions[ + int(lengths[:lo].sum()) : int(lengths[:hi].sum()) + ].tolist() + expected = { + "positions": local_pos + [0] * (3 * width - tokens), + "context_lens": (starts + lengths)[lo:hi].tolist() + + [0] * (width - hi + lo), + "cu_seqlens_q": [0] + + np.cumsum(local_lens).tolist() + + [tokens] * (width - hi + lo), + "block_tables": [[i + 1] * 4 for i in range(lo, hi)] + + [[0] * 4] * (width - hi + lo), + "v4_meta_state_slot_out": slots[lo:hi].tolist() + + [0] * (width - hi + lo), + "v4_meta_state_slot_in": slots[::-1][lo:hi].tolist() + + [0] * (width - hi + lo), + "batch_id_per_q_token": np.repeat( + np.arange(hi - lo), local_lens + ).tolist() + + [-1] * (3 * width - tokens), + } + if kv_fp8: + # Empty ubatches use the existing prefill fallback and never + # publish a decode query prefix. + expected["v4_qo_indptr"] = ( + list(range(tokens + 1)) + [tokens] * (3 * width - tokens) + if tokens + else [] + ) + for suffix, values in expected.items(): + name = f"ub{index}_{suffix}" + saved.append((outputs[runner._fv_idx][name].clone(), values)) + before = {name: var[name].cpu.clone() for name in names} + with pytest.raises(PublicationError, match="republish_reason"): + builder._prepare_ubatch_decode(**kwargs) + assert all(torch.equal(var[name].cpu, before[name]) for name in names) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for actual, values in saved: + assert actual[: len(values)].cpu().tolist() == values + assert torch.all(actual[len(values) :] == 77) + + +def test_v4_decode_rejects_borrowed_positions_before_writing(monkeypatch): + from tests.test_h2d_attention_publication import decode_batch + + runner, builder = make_runner(monkeypatch) + # The smaller indexer fixture names its common buffers separately. + runner.h2d_groups["prefill"] = runner.h2d_owner.group( + "prefill", + [ + runner.forward_vars[name]._publication + for name in ("context_lens", "block_tables") + ], + ) + runner.arange_np = np.arange(32, dtype=np.int64) + batch = decode_batch(2) + batch.num_spec_step = 0 + positions = runner.forward_vars["positions"] + tables = runner.forward_vars["block_tables"] + positions.cpu.fill_(7) + cu = runner.forward_vars["cu_seqlens_q"] + cu.cpu[:3] = torch.tensor([0, 1, 2]) + runner._gate_staging_reuse() + torch.cuda._sleep(20_000_000) + positions.copy_to_gpu(4) + tables.copy_to_gpu(2) + observed = positions.gpu[:4].clone() + try: + with pytest.raises(PublicationError, match="republish_reason"): + builder.prepare_decode(batch, 4, 4, 1) + finally: + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + assert positions.cpu[:4].tolist() == [7] * 4 + assert observed.tolist() == [7] * 4 diff --git a/tests/test_h2d_v4_publication.py b/tests/test_h2d_v4_publication.py new file mode 100644 index 0000000000..bffcba7f0b --- /dev/null +++ b/tests/test_h2d_v4_publication.py @@ -0,0 +1,419 @@ +# SPDX-License-Identifier: MIT +"""V4/V4.1 metadata producers, isolated from unsupported full-model modes.""" + +import os +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from atom.utils import CpuGpuBuffer +from atom.utils.h2d import PublicationError +from tests.attentions.deepseek_v41.helpers import ( + PagedRequest, + prepare_step, + publish_tables, +) + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1", reason="set RUN_H2D_GPU_TESTS=1" +) + + +def v41_runner(monkeypatch, transport, pp_size): + from atom.model_engine.model_runner import ModelRunner + from atom.model_ops.attentions.deepseek_v41.backend import ( + DeepseekV41MetadataBuilder, + ) + from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry + from tests.attentions.deepseek_v41.helpers import metadata_buffers + + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace(pipeline_parallel_size=pp_size) + runner.enforce_eager = True + geometry = V41PoolGeometry( + 2, ((0, 1), (1, 2)), 32, 4, 128, 32, speculative_tokens=3 + ) + runner.forward_vars = metadata_buffers(4, 32, 4, runner.device, geometry) + for name in ("positions", "batch_id_per_q_token"): + runner.forward_vars[name].publication_group = "v41_step" + runner.forward_vars["cu_seqlens_q"].publication_group = "early" + runner.forward_vars["block_tables"].publication_group = "prefill" + runner.forward_vars["decode_src"] = CpuGpuBuffer( + 4, dtype=torch.int32, device="cuda" + ) + runner.tokenID_processor = SimpleNamespace() + builder = DeepseekV41MetadataBuilder.__new__(DeepseekV41MetadataBuilder) + builder.geometry = geometry + builder.model_runner = runner + runner.attn_metadata_builder = builder + monkeypatch.setenv("ATOM_H2D_BACKEND", transport) + runner._init_forward_vars_ring() + runner._init_h2d_publication() + return runner, builder + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_v41_plans_and_step_replay_preserve_padding_and_slots( + monkeypatch, transport, pp_size +): + + runner, builder = v41_runner(monkeypatch, transport, pp_size) + graphs, copies = [], [] + for variables in runner._fv_ring: + destinations = { + name: torch.empty_like(buf.gpu) for name, buf in variables.items() + } + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for name, buf in variables.items(): + destinations[name].copy_(buf.gpu) + graphs.append(graph) + copies.append(destinations) + torch.cuda.synchronize() + saved = [] + for step_id, lengths in enumerate(([4, 2, 1], [], [1], [3, 4], [2, 1, 4], [])): + runner._advance_forward_vars() + runner._gate_staging_reuse() + variables = runner.forward_vars + for buf in variables.values(): + buf.gpu.fill_(77) + lengths = np.asarray(lengths, dtype=np.int32) + starts = np.arange(len(lengths), dtype=np.int32) * 5 + step_id + spans, offset = [], 0 + for i, (start, length) in enumerate(zip(starts, lengths)): + spans.append(PagedRequest(i, int(start), offset, int(length), i, (i,))) + offset += int(length) + torch.cuda._sleep(2_000_000) + plans = builder._build_compress_plans( + lengths, starts + lengths, running_bs=4, max_q_len=4, extra_write=3 + ) + metadata = prepare_step( + spans, + runner.device, + buffers=variables, + running_bs=4, + running_tokens=16, + max_q_len=4, + ratios=(1, 2), + state_slot_out=torch.arange(4, device="cuda", dtype=torch.int32), + publication_group=runner.h2d_groups["v41_step"], + ) + expected = {} + for group_name in ("v4_plans", "v41_step"): + group = runner.h2d_groups[group_name] + for member, count in zip(group.members, group.counts): + expected[member.name] = (member.source[:count].clone(), count) + # The second invocation must fail BEFORE it can overwrite a live source. + before = {n: variables[n].cpu.clone() for n in expected} + with pytest.raises(PublicationError, match="republish_reason"): + builder._build_compress_plans( + lengths, starts + lengths + 100, extra_write=0 + ) + with pytest.raises(PublicationError, match="republish_reason"): + prepare_step( + spans, + runner.device, + buffers=variables, + ratios=(1, 2), + publication_group=runner.h2d_groups["v41_step"], + ) + assert all(torch.equal(variables[n].cpu, value) for n, value in before.items()) + graphs[runner._fv_idx].replay() + saved.append( + ( + {n: copies[runner._fv_idx][n].clone() for n in expected}, + expected, + metadata.positions.clone(), + metadata.batch_ids.clone(), + spans, + ) + ) + for plan in plans.values(): + assert plan.write_plan_gpu.shape[0] == 16 + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for actual, expected, positions, batch_ids, spans in saved: + for name, (values, count) in expected.items(): + assert torch.equal( + actual[name][:count].cpu().view(torch.uint8), values.view(torch.uint8) + ) + assert torch.all(actual[name][count:].cpu() == 77) + expected_positions = [ + p for span in spans for p in range(span.position, span.end) + ] + expected_ids = [i for i, span in enumerate(spans) for _ in range(span.length)] + n = len(expected_positions) + assert positions.cpu().tolist() == expected_positions + [0] * (16 - n) + assert batch_ids.cpu().tolist() == expected_ids + [-1] * (16 - n) + + +@pytest.mark.parametrize("transport", ["direct", "packed"]) +def test_v4_tbo_plan_allocations_follow_actual_runner_slots(monkeypatch, transport): + from atom.model_engine.model_runner import ModelRunner + from atom.model_ops.attentions.deepseek_v4_attn import ( + DeepseekV4AttentionMetadataBuilder, + ) + + runner = ModelRunner.__new__(ModelRunner) + runner.device = torch.device("cuda", 0) + runner.config = SimpleNamespace( + pipeline_parallel_size=2, enable_tbo=True, enable_tbo_decode=True + ) + runner.enforce_eager = True + runner.forward_vars = { + name: CpuGpuBuffer(32, dtype=torch.int32, device="cuda") + for name in ("input_ids", "decode_src") + } + runner.tokenID_processor = SimpleNamespace() + builder = DeepseekV4AttentionMetadataBuilder.__new__( + DeepseekV4AttentionMetadataBuilder + ) + builder.device = runner.device + builder.model_runner = runner + builder.max_num_batched_tokens = 32 + builder.max_bs = 4 + builder.window_size = 4 + builder.max_decode_tokens = 16 + builder.index_topk = 4 + builder.max_committed_hca = 4 + builder.block_table_cols = 4 + builder._indexer_fp4 = False + builder._unique_compress_ratios_overlap = [(4, True)] + builder._alloc_v4_metadata_buffers() + monkeypatch.setenv("ATOM_H2D_BACKEND", transport) + runner._init_forward_vars_ring() + runner._init_h2d_publication() + previous = None + results = [] + for step in range(5): + runner._advance_forward_vars() + runner._gate_staging_reuse() + current = builder._get_ubatch_compress_plan_buffers(0)[4]["compress"] + assert current is runner.forward_vars["ub0_v4_compress_plan_4"] + assert current is not previous + previous = current + torch.cuda._sleep(2_000_000) + per_step = [] + for index in range(2): + context = np.array([4 + step * 4 + index], dtype=np.int32) + plans = builder._build_compress_plans( + np.array([4], dtype=np.int32), + context, + extra_write=0, + buf_prefix_ubatch=f"ub{index}_", + ) + per_step.append( + (plans[4].compress_plan_gpu.clone(), plans[4].compress_plan_cpu.copy()) + ) + results.extend(per_step) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for actual, expected in results: + np.testing.assert_array_equal(actual.cpu().numpy(), expected) + + +def test_v4_mtp_uses_immutable_prefix_after_verify_owner_is_sealed(monkeypatch): + from atom.model_engine.kv_block import STATE_SLOT_CLASS + from atom.model_ops.attentions import deepseek_v4_attn as v4 + from atom.utils.h2d import PublicationOwner + + buf = CpuGpuBuffer(9, dtype=torch.int32, device="cuda") + owner = PublicationOwner("cuda", torch.cuda.Event()) + owner.bind(buf, "v4_qo_indptr") + constants = torch.arange(9, device="cuda", dtype=torch.int32) + variables = { + "v4_qo_indptr": buf, + "v4_draft_qo_indptr": constants, + "v4_empty_kv_indptr": torch.zeros(9, device="cuda", dtype=torch.int32), + "v4_kv_indptr_swa": torch.zeros(9, device="cuda", dtype=torch.int32), + "v4_kv_indices_swa": torch.zeros(32, device="cuda", dtype=torch.int32), + "batch_id_per_q_token": CpuGpuBuffer(8, dtype=torch.int32, device="cuda"), + } + owner.completion.record() + owner.begin() + buf.np[:] = [0, 1, 2, 3, 4, 4, 4, 4, 4] + torch.cuda._sleep(2_000_000) + buf.copy_to_gpu() + owner.finish() + builder = v4.DeepseekV4AttentionMetadataBuilder.__new__( + v4.DeepseekV4AttentionMetadataBuilder + ) + builder.model_runner = SimpleNamespace( + forward_vars=variables, pool_plan=SimpleNamespace(entries={STATE_SLOT_CLASS: 2}) + ) + builder._mtp_layers_are_swa_only = True + builder.window_size = 4 + builder._kv_fp8 = True + builder.row_ids = torch.arange(8, device="cuda", dtype=torch.int32) + builder.pool_geometry = object() + builder._dest_row_buffers = dict + metadata = SimpleNamespace( + context_lens=torch.full((4,), 7, device="cuda", dtype=torch.int32), + state_slot_out=torch.arange(4, device="cuda", dtype=torch.int32), + ) + monkeypatch.setattr( + v4, "get_forward_context", lambda: SimpleNamespace(attn_metadata=metadata) + ) + monkeypatch.setattr(v4, "write_v4_paged_decode_indices", lambda **kwargs: None) + builder.prepare_mtp_decode(2, 1, 7, torch.arange(4, device="cuda")) + assert metadata.qo_indptr.data_ptr() == constants.data_ptr() + assert metadata.qo_indptr.cpu().tolist() == [0, 1, 2, 3, 4] + assert buf.gpu.cpu().tolist() == [0, 1, 2, 3, 4, 4, 4, 4, 4] + assert buf.cpu.tolist() == [0, 1, 2, 3, 4, 4, 4, 4, 4] + + +@pytest.mark.parametrize("pp_size", [1, 2]) +def test_v41_combined_publication_precedes_indptr_consumer(monkeypatch, pp_size): + from atom.model_ops.attentions.deepseek_v41 import cache as cache_module + + runner, builder = v41_runner(monkeypatch, "packed", pp_size) + builder.device = runner.device + builder.block_size = builder.geometry.block_size + builder.cache = cache_module.PagedAttentionCache( + builder.geometry, 8, 4, runner.device, max_tokens=32 + ) + original = cache_module.fill_step_indptrs + observed = [] + + def consume(step, geometry, buffers): + group = runner.h2d_groups["v41_metadata"] + assert group._backend.kernel is not None + assert runner.h2d_groups["v4_plans"]._backend is None + assert runner.h2d_groups["v41_step"]._backend is None + # Clone before the real first consumer, on the same stream: this sees + # the bytes available to indptr construction, including all padding. + observed.append( + [ + (member.destination[:count].clone(), member.source[:count].clone()) + for member, count in zip(group.members, group.counts) + if count is not None + ] + ) + return original(step, geometry, buffers) + + monkeypatch.setattr(cache_module, "fill_step_indptrs", consume) + keys = {} + for iteration, lengths in enumerate(([4, 1, 2],) * 4 + ([1, 3],) * 4): + runner._advance_forward_vars() + runner._gate_staging_reuse() + slots = [2, 0, 1][: len(lengths)] + blocks = tuple((slot,) for slot in slots) + batch = SimpleNamespace( + is_dummy_run=False, + req_ids=tuple(100 + slot for slot in slots), + num_scheduled_tokens=np.asarray(lengths, dtype=np.int32), + context_lens=np.asarray(lengths, dtype=np.int32) + 8 + iteration, + state_slots_committed=slots, + block_tables=blocks, + total_seqs_num=len(lengths), + total_tokens_num=sum(lengths), + ) + torch.cuda._sleep(2_000_000) + metadata, _ = builder._prepare(batch, 4, 16, max_q_len=4, tentative=True) + group = runner.h2d_groups["v41_metadata"] + table_count = group.counts[group.indices["block_tables"]] + assert table_count == (None if keys.get(runner._fv_idx) == blocks else 4) + keys[runner._fv_idx] = blocks + assert metadata.state_slot_out.shape == (4,) + assert metadata.step.positions.shape == (16,) + # A plan was already included in the combined publication: its own + # group still rejects duplicate writes before touching pinned sources. + with pytest.raises(PublicationError, match="republish_reason"): + builder._build_compress_plans( + batch.num_scheduled_tokens, batch.context_lens, extra_write=0 + ) + runner._mark_staging_h2d_enqueued() + runner._record_forward_vars_event() + torch.cuda.synchronize() + for fields in observed: + for gpu, cpu in fields: + assert torch.equal(gpu.cpu().view(torch.uint8), cpu.view(torch.uint8)) + + +@pytest.mark.parametrize("source", ["query_prefix", "block_tables"]) +def test_v41_standalone_producer_rejects_before_source_rewrite(monkeypatch, source): + + runner, _ = v41_runner(monkeypatch, "direct", 1) + var = runner.forward_vars + spans = [PagedRequest(0, 4, 0, 1, 0, (7,))] + runner._gate_staging_reuse() + if source == "query_prefix": + buf = var["cu_seqlens_q"] + buf.np[:3] = [0, 2, 2] + else: + buf = var["block_tables"] + buf.np[:2] = 3 + before = buf.cpu.clone() + torch.cuda._sleep(20_000_000) + buf.copy_to_gpu(2 if source == "block_tables" else 3) + observed = buf.gpu.clone() + try: + with pytest.raises(PublicationError, match="republish_reason"): + if source == "block_tables": + publish_tables(buf, spans, 2) + else: + prepare_step( + spans, + runner.device, + buffers=var, + running_bs=2, + running_tokens=4, + state_slot_out=torch.zeros(2, dtype=torch.int32, device="cuda"), + publication_group=runner.h2d_groups["v41_step"], + ) + finally: + runner._mark_staging_h2d_enqueued() + torch.cuda.synchronize() + assert torch.equal(buf.cpu, before) + count = 2 if source == "block_tables" else 3 + assert torch.equal(observed[:count].cpu(), before[:count]) + + +def test_dummy_storage_isolation_survives_shared_metadata(): + from atom.model_ops.attentions.deepseek_v41.backend import ( + DeepseekV41MetadataBuilder, + ) + from atom.model_ops.attentions.deepseek_v41.cache import PagedAttentionCache + from atom.model_ops.attentions.pool_layout.v41_pool_geometry import V41PoolGeometry + from tests.attentions.deepseek_v41.helpers import metadata_buffers + + builder = DeepseekV41MetadataBuilder.__new__(DeepseekV41MetadataBuilder) + builder.geometry = V41PoolGeometry(1, ((0, 2),), 32, 4, 512, 32) + builder.block_size, builder.device = 32, "cuda" + builder.model_runner = SimpleNamespace( + forward_vars=metadata_buffers(4, 36, 4, "cuda", builder.geometry) + ) + serving = builder.cache = PagedAttentionCache(builder.geometry, 8, 4, "cuda") + serving.backing.fill_(17) + before = serving.backing.clone() + batch = SimpleNamespace( + is_dummy_run=True, + req_ids=(-1, -2), + num_scheduled_tokens=(34, 2), + context_lens=(34, 2), + state_slots_committed=(), + block_tables=((0,), (0,)), + total_seqs_num=2, + total_tokens_num=36, + ) + metadata, _ = builder._prepare(batch, 4, 36) + private, step = metadata.cache, metadata.step + assert private is not serving and metadata.dummy + assert not hasattr(step.requests[0], "block_ids") + assert step.block_tables[:2, :2].tolist() == [[0, 1], [2, 0]] + kv = torch.full((1, step.width, 512), 3, dtype=torch.bfloat16, device="cuda") + old_state, old_pages = private.state_bytes.clone(), private.page_bytes.clone() + private.write_window(0, kv, step) + n = step.plans[2].compress_plan_gpu.shape[0] + compressed = torch.full((1, n, 512), 5, dtype=torch.bfloat16, device="cuda") + private._scatter_rows(private.pages.view("main_0")[0], step, compressed, 2) + torch.cuda.synchronize() + assert torch.equal(serving.backing, before) + assert not torch.equal(private.state_bytes, old_state) + assert not torch.equal(private.page_bytes, old_pages) diff --git a/tests/test_model_runner_decode_padding.py b/tests/test_model_runner_decode_padding.py new file mode 100644 index 0000000000..7450655bab --- /dev/null +++ b/tests/test_model_runner_decode_padding.py @@ -0,0 +1,171 @@ +# SPDX-License-Identifier: MIT +"""Model execution must retain the step's padded height until sampling.""" + +from types import SimpleNamespace + +import pytest +import torch + +pytest.importorskip("aiter") + +from atom.model_engine import model_runner as runner_module +from atom.utils.forward_context import ForwardMode + + +@pytest.mark.parametrize("mrope", [False, True]) +@pytest.mark.parametrize("dummy", [False, True]) +@pytest.mark.parametrize( + "scheduled,running,seqs,q,unified,prefill,piecewise", + [ + (3, 4, 3, 1, True, False, False), + (9, 12, 3, 3, True, False, False), + (4, 9, 2, 3, True, False, False), # packed/ragged, not seqs * q + (3, 3, 3, 1, False, False, False), # a peer is prefilling + (7, 7, 2, 4, False, True, False), + (4, 4, 4, 1, True, False, False), + (3, 4, 3, 1, True, False, True), + ], +) +def test_runner_model_height_and_sampling_rows( + monkeypatch, mrope, dummy, scheduled, running, seqs, q, unified, prefill, piecewise +): + runner = runner_module.ModelRunner.__new__(runner_module.ModelRunner) + runner.use_mrope = mrope + runner.config = SimpleNamespace(prefill_context_parallel_size=1) + runner._detailed_label_suffix = lambda batch: "" + runner._piecewise_cg_active = lambda: piecewise + ids = torch.full((32,), -99, dtype=torch.int32) + ids[:scheduled] = torch.arange(1, scheduled + 1) + pos = torch.full((32,), -99, dtype=torch.int64) + pos[:scheduled] = torch.arange(11, scheduled + 11) + runner.forward_vars = { + "input_ids": SimpleNamespace(gpu=ids), + "positions": SimpleNamespace(gpu=pos), + "mrope_positions": SimpleNamespace(gpu=torch.full((96,), -99)), + } + if mrope: + # The mRoPE builder packs three planes at the running-token stride. + positions = runner._mrope_positions_view(running) + positions[:, :scheduled] = pos[:scheduled] + else: + positions = pos[:scheduled] + mode = ForwardMode( + use_cudagraph=piecewise, + is_prefill=prefill, + scheduled_bs=seqs, + scheduled_tokens=scheduled, + running_bs=4, + running_tokens=running, + running_tokens_are_unified=unified, + max_seqlen_q=q, + piecewise_captured=piecewise, + tbo_collective_active=False, + ) + metadata = SimpleNamespace( + slot_mapping=torch.full((running,), -1), + cu_seqlens_q=torch.arange(5), + context_lens=torch.tensor([8, 8, 8, 0]), + ) + original_metadata = {k: (v, v.clone()) for k, v in vars(metadata).items()} + ctx = SimpleNamespace( + context=SimpleNamespace( + scheduled_bs=seqs, + scheduled_tokens=scheduled, + running_bs=4, + is_prefill=prefill, + is_dummy_run=dummy, + positions=positions, + forward_mode=mode, + ), + attn_metadata=metadata, + ubatch_slices=None, + ) + monkeypatch.setattr(runner_module, "get_forward_context", lambda: ctx) + monkeypatch.setattr( + runner_module, "get_pp_group", lambda: SimpleNamespace(world_size=1) + ) + expected_rows = running if unified and not prefill else scheduled + + class Model: + def __call__(self, input_ids, model_positions): + assert input_ids.shape[0] == expected_rows + assert model_positions.shape[-1] == expected_rows + torch.testing.assert_close(input_ids[:scheduled], ids[:scheduled]) + if expected_rows > scheduled: + assert torch.all(input_ids[scheduled:] == 0) + assert torch.all(model_positions[..., scheduled:] == 0) + return input_ids[:, None].float() + + def compute_logits(self, hidden): + assert hidden.shape[0] == scheduled + return hidden + 1 + + runner.model = Model() + logits, hidden = runner.run_model(ids[:scheduled]) + assert hidden.shape == (scheduled, 1) + torch.testing.assert_close(logits[:, 0], torch.arange(2, scheduled + 2).float()) + assert torch.all(ids[running:] == -99) + for name, (original, values) in original_metadata.items(): + assert getattr(metadata, name) is original + torch.testing.assert_close(original, values) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires GPU") +@pytest.mark.parametrize("capture", [False, True]) +@pytest.mark.parametrize( + "query_lens", [(1,), (1, 1), (1, 1, 1), (3,), (3, 3, 3), (1, 3)] +) +def test_pa_preserves_complete_padded_layout(query_lens, capture): + """Zero-context rows belong to the output; the guard starts after them.""" + import aiter + + from atom.model_ops.base_attention import run_pa_fwd_asm + + padded = 4 + scheduled = sum(query_lens) + max_qlen = max(query_lens) + running = 9 if query_lens == (1, 3) else padded * max_qlen + q = torch.zeros((running, 32, 128), device="cuda", dtype=torch.bfloat16) + k = torch.ones((32, 8, 8, 16, 16), device="cuda", dtype=aiter.dtypes.fp8) + v = torch.ones((32, 8, 1, 128, 16), device="cuda", dtype=aiter.dtypes.fp8) + scale = torch.ones((32, 8, 16), device="cuda") + blocks = torch.arange(32, device="cuda", dtype=torch.int32).repeat(padded, 1) + lengths = torch.tensor( + [284] * len(query_lens) + [0] * (padded - len(query_lens)), + device="cuda", + dtype=torch.int32, + ) + cu = torch.tensor( + [0, *query_lens, *([0] * (padded - len(query_lens)))], + device="cuda", + dtype=torch.int32, + ).cumsum(0, dtype=torch.int32) + storage = torch.full((running + 1, 32, 128), 123, device="cuda", dtype=q.dtype) + output = storage[:running] + + def run(): + return run_pa_fwd_asm( + q, + k, + v, + blocks, + lengths, + scale, + scale, + out=output, + qo_indptr=cu, + max_qlen=max_qlen, + ) + + run() + if capture: + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + run() + storage.fill_(123) + graph.replay() + torch.cuda.synchronize() + torch.testing.assert_close(output[:scheduled], torch.ones_like(output[:scheduled])) + assert torch.all(storage[running:] == 123) + assert blocks.shape[0] == lengths.shape[0] == padded + assert cu.shape[0] == padded + 1 diff --git a/tests/test_mtp_deferred_status_queue.py b/tests/test_mtp_deferred_status_queue.py index abc8234079..737dcf0c89 100644 --- a/tests/test_mtp_deferred_status_queue.py +++ b/tests/test_mtp_deferred_status_queue.py @@ -96,6 +96,7 @@ def test_a_real_batch_still_carries_over_from_a_real_batch(): def _processor() -> tokenIDProcessor: processor = object.__new__(tokenIDProcessor) + processor.runner = SimpleNamespace(h2d_groups={"input_ids": mock.Mock()}) processor.input_ids = SimpleNamespace( np=np.zeros(8, dtype=np.int32), gpu=np.zeros(8, dtype=np.int32), diff --git a/tests/test_packed_h2d.py b/tests/test_packed_h2d.py new file mode 100644 index 0000000000..87d95f310c --- /dev/null +++ b/tests/test_packed_h2d.py @@ -0,0 +1,231 @@ +# SPDX-License-Identifier: MIT +"""Real DMA/scatter correctness, asynchronous source reuse.""" + +import os + +import pytest +import torch + +from atom.utils.h2d import PublicationError, PublicationOwner +from tests.test_h2d_publication import buffer, setup + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_H2D_GPU_TESTS") != "1", reason="set RUN_H2D_GPU_TESTS=1" +) + + +def test_mixed_bytes_padding_dynamic_counts_and_graph_replay(): + owner = PublicationOwner("cuda", torch.cuda.Event()) + layouts = [ + ((5, 3), torch.float32, "rows", 2), + ((1031,), torch.int64, "elements", 1027), + ((19,), torch.bool, "bytes", 17), + ((21,), torch.bfloat16, "bytes", 41), + ] + buffers = [ + buffer(*shape, dtype=dtype, device="cuda") for shape, dtype, _, _ in layouts + ] + members = [ + owner.bind(b, str(i), unit=layout[2]) + for i, (b, layout) in enumerate(zip(buffers, layouts)) + ] + group = owner.group("mixed", members) + assert group.use_transport("packed") == "packed" + pointers = [b.gpu.data_ptr() for b in buffers] + expected = [ + torch.full((b.cpu.numel() * b.cpu.element_size(),), 199, dtype=torch.uint8) + for b in buffers + ] + owner.begin() + group.publish([b.capacity for b in members]) # Compile before delayed work. + owner.finish() + for b in buffers: + b.gpu.view(torch.uint8).fill_(199) + outputs = [torch.empty_like(b.gpu.view(torch.uint8)) for b in buffers] + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for output, b in zip(outputs, buffers): + output.copy_(b.gpu.view(torch.uint8)) + for step, counts in enumerate( + ( + [x[3] for x in layouts], + [0, 3, None, 1], + [1, None, None, None], + [None] * 4, + [b.capacity for b in members], + ) + ): + owner.begin() + for i, b in enumerate(buffers): + b.cpu.view(torch.uint8).fill_(71 + step + i) + torch.cuda._sleep(2_000_000) + group.publish(counts) + owner.finish() + graph.replay() + # begin() in the next step must wait before the same arena is packed. + owner.completion.synchronize() + torch.cuda.synchronize() + for i, (b, member, count, output) in enumerate( + zip(buffers, members, counts, outputs) + ): + expected[i][: (count or 0) * member.bytes_per_count] = 71 + step + i + assert torch.equal(output.cpu().flatten(), expected[i]) + assert b.gpu.data_ptr() == pointers[i] + assert len(group._backend.prefixes) <= 4 + + +def test_packed_arena_reuse_and_atomic_validation(): + owner, a, b, _, _, group = setup("cuda", shape=(1031,)) + assert group.use_transport("packed") == "packed" + owner.begin() + group.publish((1031, 1031)) + owner.finish() + observations = [] + for step in range(6): + owner.begin() + a.cpu.fill_(step + 11) + b.cpu.fill_(step + 21) + header = group._backend.header.copy() + with pytest.raises(ValueError): + group.publish((17, 1032)) + assert (header == group._backend.header).all() + torch.cuda._sleep(3_000_000) + group.publish((1031, 1031)) + observations.append((a.gpu.clone(), b.gpu.clone())) + with pytest.raises(PublicationError, match="republish_reason"): + group.publish((1, None)) + owner.finish() + owner.completion.synchronize() + for step, (first, second) in enumerate(observations): + assert torch.all(first.cpu() == step + 11) + assert torch.all(second.cpu() == step + 21) + + +def test_packed_failure_after_dma_poison_retains_arena(monkeypatch): + owner, a, b, _, _, group = setup("cuda") + assert group.use_transport("packed") == "packed" + owner.begin() + a.cpu.fill_(17) + b.cpu.fill_(23) + + class FailingLaunch: + def __getitem__(self, grid): + def fail(*args, **kwargs): + raise RuntimeError("scatter enqueue failed") + + return fail + + monkeypatch.setattr(group._backend, "launch", FailingLaunch()) + with pytest.raises(RuntimeError, match="scatter enqueue"): + group.publish((16, 16)) + assert group._backend.host.is_pinned() + owner.drain() + with pytest.raises(PublicationError, match="failed"): + owner.begin() + + +def test_packed_initializes_with_cuda_as_default_device(): + with torch.device("cuda"): + owner, a, b, _, _, group = setup("cuda") + assert group.use_transport("packed") == "packed" + assert group._backend.host.device.type == "cpu" + owner.begin() + a.cpu.fill_(51) + b.cpu.fill_(61) + group.publish((16, 16)) + owner.finish() + owner.completion.synchronize() + assert torch.all(a.gpu.cpu() == 51) and torch.all(b.gpu.cpu() == 61) + + +@pytest.mark.parametrize("strided", [False, True]) +def test_owner_packs_consumer_boundaries_and_keeps_supported_fallback(strided): + owner, a, b, x, y, producer = setup("cuda") + c = buffer(16, device="cuda") + if strided: + c.cpu, c.gpu = c.cpu[::2], c.gpu[::2] + z = owner.bind(c, "c") + consumer = owner.group("consumer", (x, y, z)) + owner.use_packed_transport() + # Unsupported large groups must not disable packing a supported producer. + assert consumer.transport == ("direct" if strided else "packed") + assert producer.transport == ("packed" if strided else "direct") + assert (producer._backend is None) != (consumer._backend is None) + owner.begin() + a.cpu.fill_(11) + b.cpu.fill_(22) + c.cpu.fill_(33) + consumer.publish((16, 16, c.cpu.shape[0])) + result = (a.gpu.clone(), b.gpu.clone(), c.gpu.clone()) + with pytest.raises(PublicationError, match="republish_reason"): + producer.publish((16, 16)) + owner.finish() + # The producer remains usable independently in the next epoch. + owner.begin() + a.cpu.fill_(44) + b.cpu.fill_(55) + producer.publish((16, 16)) + owner.finish() + owner.completion.synchronize() + assert all(torch.all(t.cpu() == v) for t, v in zip(result, (11, 22, 33))) + assert torch.all(a.gpu.cpu() == 44) and torch.all(b.gpu.cpu() == 55) + + +def test_disjoint_single_member_publications_do_not_borrow_packed_header(monkeypatch): + owner, a, b, _, _, group = setup("cuda", shape=(8,)) + assert group.use_transport("packed") == "packed" + + def unexpected_wait(): + pytest.fail("a fresh member must not wait for another member's direct DMA") + + monkeypatch.setattr(owner, "_wait_sources", unexpected_wait) + owner.begin() + try: + a.cpu.fill_(17) + torch.cuda._sleep(3_000_000) + group.publish((8, None)) + first = a.gpu.clone() + b.cpu.fill_(23) + group.publish((None, 8)) + second = b.gpu.clone() + with pytest.raises(PublicationError, match="republish_reason"): + group.publish((8, None)) + finally: + owner.finish() + owner.completion.synchronize() + assert first.cpu().tolist() == [17] * 8 + assert second.cpu().tolist() == [23] * 8 + + +def test_direct_member_preserves_an_in_flight_packed_arena(): + owner = PublicationOwner("cuda", torch.cuda.Event()) + buffers = [buffer(8, device="cuda") for _ in range(5)] + members = [owner.bind(buf, str(i)) for i, buf in enumerate(buffers)] + group = owner.group("sparse", members) + assert group.use_transport("packed") == "packed" + owner.begin() + group.publish((8, 8, None, None, None)) # Compile before delaying the queue. + owner.finish() + owner.begin() + try: + for i, buf in enumerate(buffers): + buf.cpu.fill_(11 + i) + torch.cuda._sleep(3_000_000) + group.publish((8, 8, None, None, None)) + observed = [buf.gpu.clone() for buf in buffers[:2]] + # Direct DMA can publish a fresh member while the arena is borrowed. + group.publish((None, None, 8, None, None)) + observed.append(buffers[2].gpu.clone()) + # That direct copy must not release the preceding packed source. + with pytest.raises(PublicationError, match="transport counts are in flight"): + group.publish((None, None, None, 8, 8)) + members[0].acquire_write(republish_reason="finish the earlier packed read") + group.publish((None, None, None, 8, 8)) + observed.extend(buf.gpu.clone() for buf in buffers[3:]) + finally: + owner.finish() + owner.completion.synchronize() + assert [value.cpu().tolist() for value in observed] == [ + [11 + i] * 8 for i in range(5) + ] diff --git a/tests/test_qwen4_exp_mtp.py b/tests/test_qwen4_exp_mtp.py index 3f4698d705..4923ef819b 100644 --- a/tests/test_qwen4_exp_mtp.py +++ b/tests/test_qwen4_exp_mtp.py @@ -205,6 +205,7 @@ def test_decode_mrope_storage_matches_padded_graph_stride(monkeypatch, k): gpu = torch.full_like(cpu, -777) builder.model_runner = SimpleNamespace( config=SimpleNamespace(max_model_len=8192), + use_mrope=True, forward_vars={ "mrope_positions": SimpleNamespace(cpu=cpu, np=cpu.numpy(), gpu=gpu) }, @@ -215,19 +216,27 @@ def test_decode_mrope_storage_matches_padded_graph_stride(monkeypatch, k): width = k + 1 def prepare(self, batch, running_bs, running_tokens, max_seqlen_q): - real = batch.total_tokens_num_decode - self._mrope_cpu_view(real)[:] = batch.positions return ( SimpleNamespace(max_seqlen_q=max_seqlen_q), - self._copy_mrope_to_gpu(real), + self._build_mrope_decode_positions( + batch, batch.context_lens, max_seqlen_q, running_tokens=running_tokens + ), ) monkeypatch.setattr(GDNAttentionMetadataBuilder, "prepare_decode", prepare) # Shrinking requests must not expose the previous batch's axis/tail data. for scheduled in (4, 3, 1): tokens = scheduled * width - expected = np.arange(3 * tokens).reshape(3, tokens) + 100 - batch = SimpleNamespace(total_tokens_num_decode=tokens, positions=expected) + ends = np.arange(scheduled) * 100 + width + expected = np.tile( + np.concatenate([np.arange(end - width, end) for end in ends]), (3, 1) + ) + batch = SimpleNamespace( + total_tokens_num_decode=tokens, + context_lens=ends, + req_ids=tuple(range(scheduled)), + mrope_position_deltas={}, + ) _, positions = builder.prepare_decode(batch, 4, 4 * width, width) reference = torch.from_numpy(expected) torch.testing.assert_close(positions, reference) @@ -270,7 +279,10 @@ def test_draft_allocation_profile_does_not_require_qsa_pools(): slots = torch.zeros(16, dtype=torch.int64) builder = object.__new__(Qwen4ExpMetadataBuilder) builder.model_runner = SimpleNamespace( - forward_vars={"slot_mapping": SimpleNamespace(gpu=slots)} + forward_vars={ + "slot_mapping": SimpleNamespace(gpu=slots), + "context_lens": SimpleNamespace(gpu=torch.zeros(16, dtype=torch.int32)), + } ) for positions in (torch.zeros(4), torch.zeros(3, 4)): metadata = builder.prepare_mtp_decode(4, 1, 16, positions) diff --git a/tests/test_sampler_greedy_rows.py b/tests/test_sampler_greedy_rows.py index ac998f1a5b..559f3f99e1 100644 --- a/tests/test_sampler_greedy_rows.py +++ b/tests/test_sampler_greedy_rows.py @@ -1,15 +1,5 @@ -"""`_greedy_tokens` picks the right rows, whichever reducer answers. - -The greedy fixup used to reduce `probs[greedy_mask]` -- a gather that -materializes a `[greedy, vocab]` copy before reducing anything. It now reduces -every row and keeps the selected answers, which is the same result only if the -row bookkeeping lines up: the gather produced answers in mask order, and so must -indexing the full answer. - -That bookkeeping is ATOM's half and is what this covers. The reducer's half -- -that a per-row argmax equals `torch.argmax`, ties included -- belongs to aiter's -own op tests and needs a GPU, which this suite does not have. -""" +# SPDX-License-Identifier: MIT +"""Greedy correction preserves rows without a device-to-host mask decision.""" import pytest import torch @@ -21,73 +11,82 @@ ) -def _torch_topk_select(input, topk, *, tie=None, **_kwargs): - """`aiter.topk_select` over torch, for a runner with no GPU. - - The reduction is aiter's and is tested where a GPU can run it. Stubbing it - leaves exactly the half this file is about. `tie="low"` is asserted rather - than honoured -- `torch.argmax` already breaks ties that way, so a stub that - silently accepted any `tie` would let the call site lose the promise without - a test noticing. - """ - assert tie == "low", f"the sampler must ask for the lowest-index tie, not {tie!r}" - return None, torch.topk(input, topk, dim=-1, sorted=True).indices.to(torch.int32) - - -@pytest.fixture(autouse=True) -def _stub_selector(monkeypatch): - monkeypatch.setattr(sampler, "topk_select", _torch_topk_select) - - -def _probs(rows: int, vocab: int, seed: int) -> torch.Tensor: - return torch.rand(rows, vocab, generator=torch.Generator().manual_seed(seed)) +def _torch_topk_select(probs, k, *, tie): + assert k == 1 and tie == "low" + return None, probs.argmax(-1, keepdim=True).to(torch.int32) @pytest.mark.parametrize("rows,vocab", [(1, 8), (4, 16), (17, 61), (64, 129)]) @pytest.mark.parametrize("frac", [0.0, 0.25, 0.5, 1.0]) -def test_greedy_rows_match_the_gather_they_replaced(rows, vocab, frac): - probs = _probs(rows, vocab, rows * 31 + vocab) - mask = torch.zeros(rows, dtype=torch.bool) - mask[: int(rows * frac)] = True - - got = sampler._greedy_tokens(probs, mask) - - assert got.shape == (int(rows * frac),) - assert torch.equal(got, probs[mask].argmax(dim=-1)) - - -def test_greedy_rows_follow_a_scattered_mask(): - """Not just a prefix: the rows kept must be the rows the mask names. - - A contiguous mask cannot tell "index the full answer" apart from "reduce the - first N rows", which is the way this rewrite could have gone wrong. - """ - probs = _probs(9, 32, 7) - mask = torch.tensor([False, True, False, False, True, True, False, False, True]) - - got = sampler._greedy_tokens(probs, mask) - - assert torch.equal(got, probs[mask].argmax(dim=-1)) - assert torch.equal( - got, torch.stack([probs[r].argmax() for r in (1, 4, 5, 8)]).to(got.dtype) +@pytest.mark.parametrize("column", [False, True]) +def test_greedy_correction_preserves_sampled_rows( + monkeypatch, rows, vocab, frac, column +): + monkeypatch.setattr(sampler, "topk_select", _torch_topk_select) + probs = torch.rand(rows, vocab, generator=torch.Generator().manual_seed(31)) + temperatures = torch.ones(rows) + # Scattered zero-temperature rows, not just a prefix. + temperatures[torch.randperm(rows)[: int(rows * frac)]] = 0 + sampled = torch.full((rows, 1) if column else (rows,), -1, dtype=torch.long) + got = sampler._apply_greedy_tokens(probs, temperatures, sampled) + expected = torch.tensor( + [ + int(probs[row].argmax()) if temperatures[row] == 0 else -1 + for row in range(rows) + ], + dtype=torch.int32, + ) + assert got.shape == (rows,) and got.dtype == torch.int32 + assert torch.equal(got, expected) + assert (sampled == -1).all() # Caller-owned sampled results are not mutated. + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="ROCm GPU required") +@pytest.mark.parametrize("path", ["aiter", "native"]) +@pytest.mark.parametrize("greedy", ["none", "mixed", "all"]) +def test_gpu_sampling_greedy_rows_and_no_mask_readback(path, greedy): + # Real reducers and samplers. One-hot stochastic rows make their expected + # samples exact without depending on RNG state; greedy rows have ties. + rows, vocab = 8, 4096 + probs = torch.zeros(rows, vocab, device="cuda") + probs[:, 7] = 1 + temperatures = torch.ones(rows, device="cuda") + indices = {"none": [], "mixed": [1, 4, 6], "all": list(range(rows))}[greedy] + for row in indices: + temperatures[row] = 0 + probs[row, 3] = 1 + probs /= probs.sum(-1, keepdim=True) + expected = torch.full((rows,), 7, dtype=torch.int32, device="cuda") + for row in indices: + expected[row] = 3 + top_ps = torch.ones(rows, device="cuda") # Avoid independent scalar .item(). + instance = sampler.Sampler() + + def run(): + if path == "aiter": + return instance._aiter_sample( + probs, None, top_ps, False, True, temperatures + ) + return instance._native_sample(probs, None, top_ps, temperatures) + + run() # Compile before measuring. + torch.cuda.synchronize() + with torch.profiler.profile( + activities=[ + torch.profiler.ProfilerActivity.CPU, + torch.profiler.ProfilerActivity.CUDA, + ] + ) as profile: + result = run() + assert result.shape == (rows,) and result.dtype == torch.int32 + assert torch.equal(result, expected) + names = {event.key for event in profile.key_averages()} + assert not names.intersection({"aten::is_nonzero", "aten::nonzero"}) + # RNG/sampling internals can extract their own scalar values. The greedy + # correction itself must not extract any scalar, even for an empty mask. + with torch.profiler.profile() as correction_profile: + sampler._apply_greedy_tokens(probs, temperatures, result) + correction_names = {event.key for event in correction_profile.key_averages()} + assert not correction_names.intersection( + {"aten::is_nonzero", "aten::nonzero", "aten::item", "aten::_local_scalar_dense"} ) - - -def test_greedy_rows_are_written_where_the_mask_points(): - """The assignment the callers make, which is where a row mix-up would land.""" - probs = _probs(6, 24, 3) - mask = torch.tensor([False, True, True, False, False, True]) - next_tokens = torch.full((6,), -1, dtype=torch.long) - - next_tokens[mask] = sampler._greedy_tokens(probs, mask).to(next_tokens.dtype) - - for row in range(6): - expected = int(probs[row].argmax()) if mask[row] else -1 - assert int(next_tokens[row]) == expected, f"row {row}" - - -def test_an_all_false_mask_selects_nothing(): - probs = _probs(5, 12, 11) - mask = torch.zeros(5, dtype=torch.bool) - - assert sampler._greedy_tokens(probs, mask).numel() == 0 diff --git a/tests/test_sampler_scalar_filters.py b/tests/test_sampler_scalar_filters.py new file mode 100644 index 0000000000..09f46989a3 --- /dev/null +++ b/tests/test_sampler_scalar_filters.py @@ -0,0 +1,86 @@ +# SPDX-License-Identifier: MIT +"""CPU-known filters keep AITER scalar dispatch without device readback.""" + +import pytest +import torch + +sampler = pytest.importorskip("atom.model_ops.sampler", exc_type=ImportError) + + +@pytest.mark.parametrize("top_k,top_p", [(3, None), (None, 0.75), (3, 0.75)]) +def test_uniform_verification_filters_preserve_scalars(monkeypatch, top_k, top_p): + def capture(self, logits, temperatures, top_ks, top_ps, **kwargs): + assert temperatures.tolist() == [1, 1, 2, 2, 2] + assert top_ks == top_k and top_ps == top_p + assert kwargs["needs_independent_noise"] + return torch.zeros(5, dtype=torch.int32) + + monkeypatch.setattr(sampler.Sampler, "forward", capture) + sampler.Sampler().sample_verification_tokens( + torch.zeros(5, 8), + torch.tensor([0, 2, 2, 5]), + torch.tensor([0.5, 1.0, 1.5, 2.0]), + top_k, + top_p, + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="ROCm GPU required") +@pytest.mark.parametrize("path", ["aiter", "native"]) +@pytest.mark.parametrize("rows", [1, 8]) +@pytest.mark.parametrize("top_k,top_p", [(3, None), (None, 0.75), (3, 0.75), (-1, 1.0)]) +def test_scalar_filters_match_legacy_tensor_sampling( + monkeypatch, path, rows, top_k, top_p +): + monkeypatch.setattr(sampler, "AITER_TOPK_TOPP_AVAILABLE", path == "aiter") + instance = sampler.Sampler() + logits = torch.randn(rows, 4096, device="cuda") + temperatures = torch.linspace(0.5, 1.5, rows, device="cuda") + k = ( + None + if top_k is None + else torch.tensor([top_k], dtype=torch.int32, device="cuda") + ) + p = None if top_p is None else torch.tensor([top_p], device="cuda") + torch.manual_seed(71) + expected = instance(logits, temperatures, k, p) + torch.manual_seed(71) + # The real sampler must accept CPU filters without a Python tensor readback. + # Its RNG may use internal C++ scalar operations independently of filters. + with monkeypatch.context() as patch: + patch.setattr(torch.Tensor, "item", lambda *_: pytest.fail("filter readback")) + actual = instance(logits, temperatures, top_k, top_p) + assert actual.dtype == torch.int32 and actual.shape == (rows,) + assert torch.equal(actual, expected) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="ROCm GPU required") +@pytest.mark.parametrize("transport", ["direct", "packed"]) +def test_runner_uniform_filters_reach_sampler_without_readback(monkeypatch, transport): + from types import SimpleNamespace + + import numpy as np + + from tests.test_h2d_runner_publication import runner_with_buffers + + runner = runner_with_buffers(monkeypatch, transport) + batch = SimpleNamespace( + total_seqs_num=3, + temperatures=np.ones(3, dtype=np.float32), + top_ks=np.full(3, 3, dtype=np.int32), + top_ps=np.full(3, 0.75, dtype=np.float32), + ) + runner._gate_staging_reuse() + group = runner.h2d_groups.get("token_inputs") + temperatures, k, p, greedy, noise = runner.prepare_sample( + batch, publication_group=group + ) + if group is not None: + group.publish(group.counts) + runner._mark_staging_h2d_enqueued() + logits = torch.full((3, 4096), -1000.0, device="cuda") + logits[:, 7] = 0 + with monkeypatch.context() as patch: + patch.setattr(torch.Tensor, "item", lambda *_: pytest.fail("filter readback")) + actual = sampler.Sampler()(logits, temperatures, k, p, greedy, noise) + assert actual.tolist() == [7, 7, 7] diff --git a/tests/test_shared_block_tables.py b/tests/test_shared_block_tables.py new file mode 100644 index 0000000000..26d350c725 --- /dev/null +++ b/tests/test_shared_block_tables.py @@ -0,0 +1,356 @@ +# SPDX-License-Identifier: MIT +"""Page-map reuse is tied to contents, publication success and physical slots.""" + +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from atom.model_engine.block_table_codec import ( + BlockTableDeltaDecoder, + BlockTableDeltaEncoder, +) +from atom.model_engine.sequence import BlockTable, new_block_table +from atom.utils import CpuGpuBuffer +from atom.utils.block_tables import block_table_state + + +def buffer(): + return CpuGpuBuffer(4, 8, dtype=torch.int32, device="cpu", pin_memory=False) + + +def test_versioned_hits_do_not_read_page_payload_or_write_sources(monkeypatch): + buf = buffer() + row = new_block_table([3, 5]) + state = block_table_state(buf).prepare([row], pad_to=4, page_limit=8) + state.publish(4) + import atom.utils.block_tables as module + + def unexpected(*args, **kwargs): + raise AssertionError("an unchanged page mapping was read or uploaded") + + monkeypatch.setattr(module, "_int32_row", unexpected) + monkeypatch.setattr(buf, "copy_to_gpu", unexpected) + buf.np.flags.writeable = False + state.prepare([row], pad_to=4, page_limit=8).publish(4) + assert buf.gpu[0, :2].tolist() == [3, 5] + + +@pytest.mark.parametrize("versioned", [True, False]) +def test_reorder_append_replace_shrink_and_padding(versioned, monkeypatch): + make = ( + new_block_table + if versioned + else lambda values: np.asarray(values, dtype=np.int32) + ) + a, b = make([3, 5]), make([2, 6]) + buf = buffer() + copies = [] + original = buf.copy_to_gpu + + def copy(count): + copies.append(count) + return original(count) + + monkeypatch.setattr(buf, "copy_to_gpu", copy) + state = block_table_state(buf) + for rows, count, dirty in ( + ([a, b], 4, True), + ([a, b], 4, False), + ([b, a], 4, True), + ([a], 4, True), + ([a], 2, True), + ([], 2, True), + ([], 2, False), + ): + before = len(copies) + state.prepare(rows, pad_to=count, page_limit=8).publish(count) + assert len(copies) - before == dirty + expected = np.zeros((count, 8), np.int32) + for i, row in enumerate(rows): + expected[i, : len(row)] = row + np.testing.assert_array_equal(buf.gpu[:count].numpy(), expected) + state.prepare([a], pad_to=4, page_limit=8).publish(4) + a[0] = 7 + state.prepare([a], pad_to=4, page_limit=8).publish(4) + assert buf.gpu[0, 0] == 7 + + +def test_decoder_carries_append_lineage_without_mutating_older_batches(): + encoder, decoder = BlockTableDeltaEncoder(), BlockTableDeltaDecoder() + row = new_block_table([3, 5]) + + def receive(): + return decoder.decode( + encoder.encode(SimpleNamespace(req_ids=[1], block_tables=[row])) + ).block_tables[0] + + first = receive() + assert isinstance(first, BlockTable) + assert receive() is first + row.append(7) + second = receive() + assert second is not first and second.version == first.version + assert list(first) == [3, 5] and list(second) == [3, 5, 7] + row[0] = 2 + third = receive() + assert third.version != second.version + assert list(second) == [3, 5, 7] + + +@pytest.mark.parametrize("appended", [[7], [6, 7]]) +def test_append_validates_only_new_ids_and_keeps_source_resizable(appended): + import array + + buf = buffer() + row = new_block_table([3, 5]) + state = block_table_state(buf).prepare([row], page_limit=8) + row.extend(appended) # No surviving buffer export may prevent this append. + # Deliberately bypass version tracking to make an already validated prefix + # unreadable as valid page ids. This is instrumentation, not a legal caller: + # append preparation must trust that prefix and only read the new suffix. + array.array.__setitem__(row, 0, -1) + state.prepare([row], page_limit=8) + assert buf.cpu[0, : len(row)].tolist() == [3, 5] + appended + array.array.__setitem__(row, 0, 3) + before = buf.cpu.clone() + row.append(8) + with pytest.raises(ValueError, match="out of range"): + state.prepare([row], page_limit=8) + assert torch.equal(buf.cpu, before) + row.pop() # Failed validation must also release every export. + state.prepare([row], page_limit=8) + with pytest.raises(ValueError, match="out of range"): + state.prepare([row], page_limit=7) + row.append(6) + + +def test_validate_whole_batch_before_writing_and_preserve_previous_snapshot(): + buf = buffer() + state = block_table_state(buf).prepare([[1], [2]], page_limit=8) + state.publish(2) + before = buf.cpu.clone() + for invalid in ([9], [1] * 9, np.array([2**32 + 1], dtype=np.int64)): + with pytest.raises(ValueError): + state.prepare([[3], invalid], page_limit=8) + assert torch.equal(buf.cpu, before) + assert torch.equal(buf.gpu[:2], before[:2]) + + +def test_failed_publish_retries_and_destination_replacement_invalidates(monkeypatch): + buf = buffer() + state = block_table_state(buf).prepare([[3]], pad_to=2) + original = buf.copy_to_gpu + + def fail(count): + raise RuntimeError("enqueue failed") + + monkeypatch.setattr(buf, "copy_to_gpu", fail) + with pytest.raises(RuntimeError, match="enqueue failed"): + state.publish(2) + assert state.published is None + monkeypatch.setattr(buf, "copy_to_gpu", original) + state.prepare([[3]], pad_to=2).publish(2) + buf.gpu = torch.full_like(buf.gpu, -1) + state.publish(2) + assert buf.gpu[0].tolist() == [3] + [0] * 7 + + +def test_cpu_preparation_is_not_gpu_publication(): + buf = buffer() + state = block_table_state(buf) + state.prepare([[3]], pad_to=2).publish(2) + state.prepare([[5]], pad_to=2) # prefill with no GPU table consumer + assert buf.gpu[0, 0] == 3 + state.prepare([[5]], pad_to=2).publish(2) + assert buf.gpu[0, 0] == 5 + + +def test_tbo_slices_reuse_row_revisions_and_recover_after_capture(monkeypatch): + source, left, right = buffer(), buffer(), buffer() + rows = [new_block_table([i + 1]) for i in range(4)] + state = block_table_state(source).prepare(rows) + for dst, start in ((left, 0), (right, 2)): + state.slice_to(dst, start, 2, pad_to=3).publish(3) + + def unexpected(): + raise AssertionError("unchanged TBO rows must not acquire a source for writing") + + with monkeypatch.context() as patch: + patch.setattr(block_table_state(left), "_acquire", unexpected) + patch.setattr(block_table_state(right), "_acquire", unexpected) + state.prepare(rows) + state.slice_to(left, 0, 2, pad_to=3).publish(3) + state.slice_to(right, 2, 2, pad_to=3).publish(3) + # Capture writes the same physical destination through the common entry. + block_table_state(left).prepare(np.zeros((3, 8), np.int32)).publish(3) + state.slice_to(left, 0, 2, pad_to=3).publish(3) + assert left.gpu[:, 0].tolist() == [1, 2, 0, 0] + assert right.gpu[:, 0].tolist() == [3, 4, 0, 0] + + +def test_slot_clones_start_with_independent_cache_state(): + original = buffer() + state = block_table_state(original).prepare([[3]], pad_to=2) + state.publish(2) + clone = original.clone() + other = block_table_state(clone) + assert other is not state and other.published is None + other.prepare([[5]], pad_to=2).publish(2) + assert original.gpu[0, 0] == 3 and clone.gpu[0, 0] == 5 + + +def test_failed_host_write_cannot_leave_a_false_empty_batch_hit(monkeypatch): + buf = buffer() + state = block_table_state(buf).prepare([new_block_table([1])], pad_to=4) + + def partial_write(*args): + buf.np[0] = 77 + raise RuntimeError("host write failed") + + with monkeypatch.context() as patch: + patch.setattr(state, "_copy_changes", partial_write) + with pytest.raises(RuntimeError, match="host write failed"): + state.prepare([new_block_table([2]), new_block_table([3])], pad_to=4) + state.prepare([], pad_to=4).publish(4) + assert buf.gpu.count_nonzero() == 0 + + +def test_group_validation_failure_does_not_publish_a_revision(): + from atom.utils.h2d import PublicationOwner + + buf, other = buffer(), buffer() + owner = PublicationOwner("cpu") + table_binding = owner.bind(buf, "block_tables") + other_binding = owner.bind(other, "other") + group = owner.group("metadata", [table_binding, other_binding]) + owner.begin() + row = new_block_table([3]) + state = block_table_state(buf).prepare([row], pad_to=2) + group.set_count(other, 5) # exceeds capacity, after the table's valid count + with pytest.raises(ValueError): + state.publish(2, group=group) + assert state.published is None and buf.gpu.count_nonzero() == 0 + group.set_count(other, 2) + state.publish(2, group=group) + owner.finish() + owner.begin() + state.prepare([row], pad_to=2).publish(2, group=group) + assert group.counts[group.indices["block_tables"]] is None + assert buf.gpu[0, 0] == 3 + owner.finish() + + +@pytest.mark.parametrize("versioned", [True, False]) +def test_full_replacement_clears_only_exposed_tails_and_padded_rows(versioned): + make = new_block_table if versioned else lambda values: np.array(values, np.int32) + buf = buffer() + buf.np[:] = -7 # No prior snapshot: even unseen row tails need initialization. + state = block_table_state(buf) + for rows, padding in ( + ([make([1] * 8), make([2] * 3)], 2), + ([make([3] * 2), make([4] * 7)], 4), + ([make([]), make([5])], 3), + ([make([6] * 8)], 4), + ): + state.prepare(rows, pad_to=padding, page_limit=8) + expected = np.zeros((padding, 8), np.int32) + for i, row in enumerate(rows): + expected[i, : len(row)] = row + np.testing.assert_array_equal(buf.np[:padding], expected) + + +def test_late_invalid_versioned_row_and_rejected_acquisition_never_write(monkeypatch): + buf = buffer() + state = block_table_state(buf).prepare([new_block_table([1])], pad_to=4) + before, keys, revision = ( + buf.cpu.clone(), + (state.versions, state.lengths), + state.revision, + ) + rows = [new_block_table([2]), new_block_table([9])] + + def reject(): + assert torch.equal(buf.cpu, before) + raise RuntimeError("source in flight") + + monkeypatch.setattr(state, "_acquire", reject) + with pytest.raises(ValueError, match="out of range"): + state.prepare(rows, pad_to=4, page_limit=8) + rows[-1][0] = 3 + with pytest.raises(RuntimeError, match="source in flight"): + state.prepare(rows, pad_to=4, page_limit=8) + assert (state.versions, state.lengths) == keys and state.revision == revision + assert torch.equal(buf.cpu, before) + for row in rows: + row.append(4) + + +def test_tbo_inherits_appends_and_tracks_unversioned_content_changes(): + source, target = buffer(), buffer() + versioned = new_block_table([1]) + unversioned = np.array([2, 3], np.int32) + state = block_table_state(source).prepare([versioned, unversioned]) + state.slice_to(target, 0, 2, pad_to=4).publish(4) + versioned.append(4) + unversioned[0] = 5 + state.prepare([versioned, unversioned]) + state.slice_to(target, 0, 2, pad_to=4).publish(4) + assert target.gpu[:2, :2].tolist() == [[1, 4], [5, 3]] + revision = block_table_state(target).revision + state.prepare([versioned, unversioned]).slice_to(target, 0, 2, pad_to=4) + assert block_table_state(target).revision == revision + + +def test_sparse_replacement_leaves_unchanged_rows_untouched(): + # Instrument the destination with sentinels to detect redundant writes. + # Actual callers must not write behind BlockTableState's snapshot. + buf = buffer() + rows = [new_block_table([i + 1] * 8) for i in range(4)] + state = block_table_state(buf).prepare(rows) + buf.np[0] = 71 + buf.np[2:] = 72 + rows[1] = new_block_table([6, 7]) + state.prepare(rows) + assert np.all(buf.np[0] == 71) + assert np.all(buf.np[2:] == 72) + assert buf.np[1].tolist() == [6, 7] + [0] * 6 + + +def test_reorder_reuses_validation_by_lineage(monkeypatch): + import atom.utils.block_tables as module + + buf = buffer() + a, b = new_block_table([1, 2]), new_block_table([3, 4]) + state = block_table_state(buf).prepare([a, b], page_limit=8) + + def unexpected(*args): + raise AssertionError("known row contents were scanned after reordering") + + with monkeypatch.context() as patch: + patch.setattr(module, "_int32_row", unexpected) + state.prepare([b, a], page_limit=8) + assert buf.np[:2, :2].tolist() == [[3, 4], [1, 2]] + before = buf.cpu.clone() + with pytest.raises(ValueError, match="out of range"): + state.prepare([b, a], page_limit=4) + assert torch.equal(buf.cpu, before) + + +@pytest.mark.parametrize("page_limit", [8, 1 << 40]) +def test_negative_page_ids_fail_before_any_write(page_limit): + buf = buffer() + state = block_table_state(buf).prepare([new_block_table([1, 2])]) + before = buf.cpu.clone() + with pytest.raises(ValueError, match="out of range"): + state.prepare([new_block_table([3, -1])], page_limit=page_limit) + assert torch.equal(buf.cpu, before) + + +def test_zero_width_rows_and_padding(): + buf = CpuGpuBuffer(4, 0, dtype=torch.int32, device="cpu", pin_memory=False) + state = block_table_state(buf) + state.prepare([new_block_table()], pad_to=4, page_limit=0) + state.prepare([], pad_to=4).publish(4) + assert buf.gpu.shape == (4, 0) From 195b0000021ab1df09e097ed55ef150c658fe15b Mon Sep 17 00:00:00 2001 From: Lingpeng Jin <103567126+valarLip@users.noreply.github.com> Date: Fri, 25 Sep 2026 03:39:55 +0000 Subject: [PATCH 2/5] profiling: add ATOM:: runner stage annotations Label staging reuse, preparation, model execution and postprocessing while the runner's torch profiler is active. Include batch token counts on forward markers and avoid record_function scopes when profiling is disabled. --- atom/model_engine/model_runner.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/atom/model_engine/model_runner.py b/atom/model_engine/model_runner.py index 1ce9ecb958..6f4ceae01a 100644 --- a/atom/model_engine/model_runner.py +++ b/atom/model_engine/model_runner.py @@ -639,6 +639,29 @@ def prepare_draft_ids( return ret +def _profile_runner_stage(method): + """Annotate runner stages only while its torch profiler is active.""" + name = method.__name__ + label = f"ATOM::{name}" + + @wraps(method) + def wrapped(self, *args, **kwargs): + if getattr(self, "profiler", None) is None: + return method(self, *args, **kwargs) + stage_label = label + if name == "forward": + batch = args[0] if args else kwargs["batch"] + stage_label += ( + f" tokens={batch.total_tokens_num}" + f" prefill={batch.total_seqs_num_prefill}" + f" decode={batch.total_seqs_num_decode}" + ) + with record_function(stage_label): + return method(self, *args, **kwargs) + + return wrapped + + class ModelRunner: def __init__(self, rank: int, config: Config): @@ -1462,6 +1485,7 @@ def _advance_forward_vars(self): self.tokenID_processor.input_ids = self.forward_vars["input_ids"] self.tokenID_processor.decode_src = self.forward_vars["decode_src"] + @_profile_runner_stage def _gate_staging_reuse(self): """Block until the previous forward's staging H2Ds have executed. @@ -1491,6 +1515,7 @@ def _gate_staging_reuse(self): """ self.h2d_owner.begin() + @_profile_runner_stage def _mark_staging_h2d_enqueued(self): """Close the window the gate above waits on. @@ -2602,6 +2627,7 @@ def prepare_sample( return temperatures, top_ks, top_ps, all_greedy, needs_independent_noise + @_profile_runner_stage def prepare_model(self, batch: ScheduledBatch): shrunk_q = self._dspark_apply_q_bucket(batch) # The step's shape, settled once. Here rather than in prepare_inputs @@ -2929,6 +2955,7 @@ def _padded_decode_inputs(self, forward_mode: ForwardMode): positions[..., scheduled:].zero_() return ids, positions + @_profile_runner_stage @record_gpu_forward def run_model( self, @@ -3158,6 +3185,7 @@ def flush_pp_send(self) -> bool: commit_pp_send_work(self._pp_pending_send) return True + @_profile_runner_stage def postprocess( self, batch: ScheduledBatch, @@ -3373,6 +3401,7 @@ def _record_kv_cache_ready(self, batch: ScheduledBatch) -> None: if callable(callback): callback(req_ids) + @_profile_runner_stage @torch.inference_mode() @with_eplb_forward_monitor def forward(self, batch: ScheduledBatch) -> ScheduledBatchOutput: @@ -4454,6 +4483,7 @@ def allocate_kv_cache(self, num_kvcache_blocks): return True return super().allocate_kv_cache(num_kvcache_blocks) + @_profile_runner_stage @torch.inference_mode() def forward(self, batch: ScheduledBatch) -> ScheduledBatchOutput: # Decode runs the model forward on a dynamically selected (optionally From 41dcec63a950059ba9c1fed666635cf11bdad089 Mon Sep 17 00:00:00 2001 From: Lingpeng Jin <103567126+valarLip@users.noreply.github.com> Date: Fri, 25 Sep 2026 03:39:55 +0000 Subject: [PATCH 3/5] perf: reuse SWA write kernels across prefill lengths Read the per-batch write width from the launch grid instead of specializing on each prefill length. Keep head_dim constexpr for the fixed layout and retain source bounds, ring indexing and padding behavior. Cover changing widths and odd head dimensions with GPU ring-write regressions. --- atom/model_ops/v4_kernels/state_writes.py | 17 ++++---- tests/test_swa_write_ring.py | 53 +++++++++++++++++++++++ 2 files changed, 61 insertions(+), 9 deletions(-) diff --git a/atom/model_ops/v4_kernels/state_writes.py b/atom/model_ops/v4_kernels/state_writes.py index b4b567cd72..66d39faa51 100644 --- a/atom/model_ops/v4_kernels/state_writes.py +++ b/atom/model_ops/v4_kernels/state_writes.py @@ -44,8 +44,8 @@ the ring size `window_size + max_spec_steps` (e.g. 128 + 0 = 128 non-MTP; 128 + 1 = 129 MTP-1). - `write_per_batch` int — max tokens to write per seq this fwd - (= `min(max_q_len, window.ring_slots)`). Used as Triton - `constexpr` for grid sizing. + (= `min(max_q_len, window.ring_slots)`). Sets grid Y; + read from the launch grid inside the kernel. Grid = `(bs, write_per_batch)`; each program writes one (seq, row-in-seq) token. Per-seq actual count is `min(token_num_per_seq[bs], write_per_batch)`; @@ -77,18 +77,17 @@ def _swa_write_kernel( pool_ptr, # this layer's whole unified-pool view, [rows, head_dim] pool_row_stride, # = head_dim pool_rows, # rows in that view; nothing may be written past it - head_dim, + head_dim: tl.constexpr, ring_start, - WRITE_PER_BATCH: tl.constexpr, BLOCK_D: tl.constexpr, RING_SLOTS: tl.constexpr, SLOT_ROWS: tl.constexpr, RING_STRIDE: tl.constexpr, RUN_ROWS: tl.constexpr, ): - """SWA ring write. 2D grid `(bs, WRITE_PER_BATCH)`. Program `(b, r)` + """SWA ring write. 2D grid `(bs, write_per_batch)`. Program `(b, r)` writes the `r`-th of the last-N tokens of seq `b`, where - `N = min(tok_n_b, WRITE_PER_BATCH)` and + `N = min(tok_n_b, grid_y)` and `tok_n_b = cu_seqlens_q[b+1] - cu_seqlens_q[b]`. Threads with `r >= N` bail. `src_id = cu_seqlens_q[b+1] - N + r` — selects directly from `kv` / @@ -112,7 +111,8 @@ def _swa_write_kernel( tok_n = cu_end - cu_start if tok_n <= 0: return - write_n = tl.minimum(tok_n, WRITE_PER_BATCH) + # The width varies with prefill length; keep it out of the JIT cache key. + write_n = tl.minimum(tok_n, tl.num_programs(1)) if row_in_batch >= write_n: return @@ -196,7 +196,7 @@ def swa_write( window: this layer's compress class's `WindowParams` (`UnifiedPoolGeometry.window_params`). write_per_batch: `min(max_q_len, window.ring_slots)` — max tokens written - per seq this fwd (grid y dim, kernel `constexpr`). + per seq this fwd (grid y dim). k_packed: [T, 512] or [T, 1, 512] fp8 NoPE extend K — fp8 2buff path only. k_rope: [T, rope_head_dim] or [T, 1, rope_head_dim] bf16 RoPE tail — fp8 2buff path only. @@ -256,7 +256,6 @@ def swa_write( pool.shape[0], head_dim, window.ring_start, - WRITE_PER_BATCH=write_per_batch, BLOCK_D=BLOCK_D, **window_constexprs(window), ) diff --git a/tests/test_swa_write_ring.py b/tests/test_swa_write_ring.py index f773c7a3aa..ea65394528 100644 --- a/tests/test_swa_write_ring.py +++ b/tests/test_swa_write_ring.py @@ -36,6 +36,7 @@ UnifiedPoolGeometry, ) from atom.model_ops.v4_kernels.state_writes import ( + _swa_write_kernel, swa_scatter_rows, swa_scatter_rows_reference, swa_write, @@ -152,6 +153,58 @@ def test_a_wrapping_seq_leaves_exactly_one_ring_live(written, ratio): assert live == RING_SLOTS, f"seq2 wrote {TOK_COUNTS[2]} tokens, {live} rows live" +@pytest.mark.parametrize("ratio", RATIO_IDS) +@pytest.mark.parametrize("capture", [False, True], ids=["eager", "graph"]) +@pytest.mark.parametrize("head_dim", [HEAD_DIM - 1, HEAD_DIM, 64, 512]) +def test_write_widths_reuse_kernel_with_padding( + geometry, batch, ratio, capture, head_dim, monkeypatch +): + """New prefill widths reuse code while retaining last-N and padded-row semantics.""" + params = geometry.window_params(ratio) + kv = torch.randn(batch["kv"].shape[0], head_dim, dtype=torch.bfloat16, device=DEV) + positions = batch["positions"].clone() + # Insert an empty request with an invalid slot between two live requests. + cu = batch["cu"][[0, 1, 1, 2, 3]] + slots = torch.tensor( + [SLOTS[0], -1, SLOTS[1], SLOTS[2]], dtype=torch.int32, device=DEV + ) + args = (kv, positions, cu, slots) + got = torch.zeros(geometry.plane_rows, head_dim, dtype=kv.dtype, device=DEV) + ref = torch.empty_like(got) + kernels = [] + launch = _swa_write_kernel.run + + def record_kernel(*args, **kwargs): + kernel = launch(*args, **kwargs) + kernels.append(kernel) + return kernel + + monkeypatch.setattr(_swa_write_kernel, "run", record_kernel) + for width in (1, 2, 6, RING_SLOTS - 1, RING_SLOTS, 3): + # Warm the launch before capture, as the engine does. + swa_write(*args, got, params, width) + graph = None + if capture: + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + swa_write(*args, got, params, width) + for _ in range(2): + # Replay must use the current device inputs at the captured addresses. + kv.add_(1) + positions.add_(RING_SLOTS + 3) + got.fill_(-1) + ref.fill_(-1) + swa_write_reference(*args, ref, params, width) + if graph is None: + swa_write(*args, got, params, width) + else: + graph.replay() + torch.cuda.synchronize() + assert torch.equal(got, ref), f"width={width}, capture={capture}" + assert kernels[0] is not None + assert all(kernel is kernels[0] for kernel in kernels), "write width triggered JIT" + + def test_over_wide_write_is_rejected(written, batch): """Must fail loudly, not race. The one contract the paged predecessor did not need: block addressing was injective on position, a ring is not.""" From ed95be6f455ce56ca72d9f45eee978e676d22b0c Mon Sep 17 00:00:00 2001 From: Lingpeng Jin <103567126+valarLip@users.noreply.github.com> Date: Fri, 25 Sep 2026 05:06:11 +0000 Subject: [PATCH 4/5] perf: reuse Triton kernels across dynamic batch and sequence lengths Move runtime bounds and strides out of exact-value specialization across Gemma RMSNorm, MRoPE, cached KV quantization, MLA conversion, Qwen4 metadata, M3 context partitioning, and V4.1 packed gather. Read redundant counts from the launch grid while retaining model geometry and useful bounded variants. Preserve vectorized KV accesses with fixed row widths and int64 lengths. Unroll only the existing small-input reduction configuration and avoid allocating unused single-pass scratch. Keep all bounds and padding guards. Add one focused suite covering output correctness, tail boundaries, large integer signatures, and compiled-kernel reuse. Relevant GPU regressions: 176 passed. Black, Ruff, and diff checks passed for this change. --- atom/model_ops/attention_mla.py | 18 +- .../attentions/deepseek_v41/packed_rows.py | 6 +- .../minimax_m3/indexer_candidate_exchange.py | 31 +- .../minimax_m3/indexer_context_parallel.py | 13 +- atom/model_ops/qwen4_exp/ops/ple.py | 3 +- atom/model_ops/qwen4_exp/ops/qsa.py | 6 +- atom/model_ops/triton_fused_qkv_quant.py | 54 ++- atom/model_ops/triton_gemma_rmsnorm.py | 6 +- atom/model_ops/triton_mrope.py | 8 +- tests/model_ops/test_triton_dynamic_shapes.py | 361 ++++++++++++++++++ 10 files changed, 439 insertions(+), 67 deletions(-) create mode 100644 tests/model_ops/test_triton_dynamic_shapes.py diff --git a/atom/model_ops/attention_mla.py b/atom/model_ops/attention_mla.py index 6b5f281298..b013ea6b05 100644 --- a/atom/model_ops/attention_mla.py +++ b/atom/model_ops/attention_mla.py @@ -3151,7 +3151,7 @@ def forward( ) -@triton.jit +@triton.jit(do_not_specialize=["OUT_NUMEL", "TOKEN_ROWS", "KV_INDICES_NUMEL"]) def _convert_req_index_to_global_index_kernel( qo_indptr, # int32 [num_requests] kv_indptr, # int32 [num_requests+1] @@ -3159,11 +3159,11 @@ def _convert_req_index_to_global_index_kernel( kv_indices, # int32 [num_requests * max_num_blocks_per_req] token_indices_ptr, # int32 [num_tokens, NUM_TOPK_TOKENS] out_kv_indices, # int32 - # shapes (compile-time where possible) + # Fixed top-k/tile geometry; runtime bounds retain all safety masks. NUM_TOPK_TOKENS: tl.constexpr, - OUT_NUMEL: tl.constexpr, - TOKEN_ROWS: tl.constexpr, - KV_INDICES_NUMEL: tl.constexpr, + OUT_NUMEL, + TOKEN_ROWS, + KV_INDICES_NUMEL, BLOCK_SIZE: tl.constexpr, BLOCK_N: tl.constexpr, # tile width along columns # strides (in elements) @@ -3321,7 +3321,7 @@ def triton_convert_req_index_to_global_index( return new_kv_indices -@triton.jit +@triton.jit(do_not_specialize=["OUT_NUMEL", "NUM_REQ"]) def _convert_req_index_to_global_index_dsa_prefill_kernel( dsa_qo_indptr, # int32 [num_tokens + 1] dsa_kv_indptr, # int32 [num_tokens + 1] @@ -3330,10 +3330,10 @@ def _convert_req_index_to_global_index_dsa_prefill_kernel( block_table, # int32 [num_req, max_num_blocks_per_req] cu_seqlens_q, # int32 [num_tokens + 1] out_kv_indices, # int32 - # shapes (compile-time where possible) + # Fixed top-k/tile geometry; runtime bounds retain all safety masks. NUM_TOPK_TOKENS: tl.constexpr, - OUT_NUMEL: tl.constexpr, - NUM_REQ: tl.constexpr, + OUT_NUMEL, + NUM_REQ, MAX_NUM_BLOCKS_PER_REQ: tl.constexpr, PAGE_SIZE: tl.constexpr, BLOCK_N: tl.constexpr, # tile width along columns diff --git a/atom/model_ops/attentions/deepseek_v41/packed_rows.py b/atom/model_ops/attentions/deepseek_v41/packed_rows.py index 44ddad5359..32bd2ab301 100644 --- a/atom/model_ops/attentions/deepseek_v41/packed_rows.py +++ b/atom/model_ops/attentions/deepseek_v41/packed_rows.py @@ -114,10 +114,8 @@ def write_packed_window(values, scales, pool, step, window, dim): ) -@triton.jit -def _gather_prefix( - pool, addresses, ptr, out, first, end, CAPACITY: tl.constexpr, D: tl.constexpr -): +@triton.jit(do_not_specialize=["first", "end", "CAPACITY"]) +def _gather_prefix(pool, addresses, ptr, out, first, end, CAPACITY, D: tl.constexpr): i = tl.program_id(0) * 16 + tl.arange(0, 16) start = tl.load(ptr + first) count = tl.load(ptr + end) - start diff --git a/atom/model_ops/minimax_m3/indexer_candidate_exchange.py b/atom/model_ops/minimax_m3/indexer_candidate_exchange.py index aafa39345a..093216f009 100644 --- a/atom/model_ops/minimax_m3/indexer_candidate_exchange.py +++ b/atom/model_ops/minimax_m3/indexer_candidate_exchange.py @@ -37,16 +37,15 @@ def _force(score, block, valid, local_start, INIT_BLOCKS: tl.constexpr): return tl.where(valid & (block >= local_start), 1e29, score) -@triton.jit +# Retain bounded alignment specialization for the score row stride. +@triton.jit(do_not_specialize=["GLOBAL_BLOCKS"]) def _local_topk( Scores, Keys, Lengths, - TOKENS: tl.constexpr, - HEADS: tl.constexpr, QUERY_LEN: tl.constexpr, - LOCAL_BLOCKS: tl.constexpr, - GLOBAL_BLOCKS: tl.constexpr, + LOCAL_BLOCKS, + GLOBAL_BLOCKS, RANK: tl.constexpr, WORLD: tl.constexpr, TOPK: tl.constexpr, @@ -68,13 +67,14 @@ def _local_topk( """ row = tl.program_id(0) head = tl.program_id(1) + tokens = tl.num_programs(0) request = row // QUERY_LEN token = row % QUERY_LEN length = tl.load(Lengths + request) # Blocks this query token may attend, in global numbering. causal_blocks = (length - QUERY_LEN + token + 128) // 128 local_start = tl.maximum(0, causal_blocks - LOCAL_KEEP) - s_row = Scores + (head * TOKENS + row) * LOCAL_BLOCKS + s_row = Scores + (head * tokens + row) * LOCAL_BLOCKS off = tl.arange(0, BLOCK_SIZE_K) local_valid = off < LOCAL_BLOCKS @@ -93,10 +93,14 @@ def _local_topk( tile = tl.topk(_pack_score_key(score, block + 1, valid), BLOCK_SIZE_T) winners = tl.topk(tl.cat(winners, tile, can_reorder=True), BLOCK_SIZE_T) off_t = tl.arange(0, BLOCK_SIZE_T) - tl.store(Keys + (head * TOKENS + row) * TOPK + off_t, winners, mask=off_t < TOPK) + tl.store( + Keys + (head * tokens + row) * TOPK + off_t, + winners, + mask=off_t < TOPK, + ) -@triton.jit +@triton.jit(do_not_specialize=["SRC_STRIDE"]) def _merge_topk( Keys, Indices, @@ -106,7 +110,6 @@ def _merge_topk( SparseCtx, TABLE_STRIDE: tl.constexpr, SBT_STRIDE: tl.constexpr, - TOKENS: tl.constexpr, QUERY_LEN: tl.constexpr, TOPK: tl.constexpr, INIT_BLOCKS: tl.constexpr, @@ -114,7 +117,7 @@ def _merge_topk( NUM_KV_HEADS: tl.constexpr, PAGES_PER_BLOCK: tl.constexpr, KEYS_PER_SHARD: tl.constexpr, - SRC_STRIDE: tl.constexpr, + SRC_STRIDE, ROW_STRIDE: tl.constexpr, REAL_CANDIDATES: tl.constexpr, CANDIDATES: tl.constexpr, @@ -221,13 +224,14 @@ def local_candidate_keys( # function is correct for the topk its own guard admits (<= 512), rather # than only for the one value production happens to pass. width = max(16, triton.next_power_of_2(local), triton.next_power_of_2(topk)) - width = min(width, 1024) + # Select medium rows in one tile; smaller streaming tiles avoid + # excessive padding work when a longer row crosses a tile boundary. + if width > 2048: + width = 1024 _local_topk[(tokens, heads)]( scores, keys, seq_lens, - TOKENS=tokens, - HEADS=heads, QUERY_LEN=max_query_len, LOCAL_BLOCKS=local, GLOBAL_BLOCKS=global_blocks, @@ -298,7 +302,6 @@ def merge_candidate_keys( output[1], TABLE_STRIDE=block_table.stride(0), SBT_STRIDE=output[0].stride(0), - TOKENS=tokens, QUERY_LEN=max_query_len, TOPK=topk, INIT_BLOCKS=init_blocks, diff --git a/atom/model_ops/minimax_m3/indexer_context_parallel.py b/atom/model_ops/minimax_m3/indexer_context_parallel.py index f7bf9361f9..5f88ad67da 100644 --- a/atom/model_ops/minimax_m3/indexer_context_parallel.py +++ b/atom/model_ops/minimax_m3/indexer_context_parallel.py @@ -14,7 +14,7 @@ ) -@triton.jit +@triton.jit(do_not_specialize=["LOCAL_BLOCKS", "GLOBAL_BLOCKS", "CHUNK"]) def _context_score( Q, Cache, @@ -24,18 +24,18 @@ def _context_score( Q_TOKEN_STRIDE: tl.constexpr, Q_HEAD_STRIDE: tl.constexpr, TABLE_STRIDE: tl.constexpr, - TOKENS: tl.constexpr, HEADS: tl.constexpr, QUERY_LEN: tl.constexpr, - LOCAL_BLOCKS: tl.constexpr, - GLOBAL_BLOCKS: tl.constexpr, + LOCAL_BLOCKS, + GLOBAL_BLOCKS, RANK: tl.constexpr, WORLD: tl.constexpr, - CHUNK: tl.constexpr, + CHUNK, N: tl.constexpr, SCALE: tl.constexpr, ): request = tl.program_id(0) + tokens = tl.num_programs(0) * QUERY_LEN chunk = tl.program_id(1) n = tl.arange(0, N) token, head = n // HEADS, n % HEADS @@ -63,7 +63,7 @@ def _context_score( ) score = tl.max(dot, 0) tl.store( - Scores + (head * TOKENS + row) * LOCAL_BLOCKS + local, + Scores + (head * tokens + row) * LOCAL_BLOCKS + local, score, mask=n < HEADS * QUERY_LEN, ) @@ -150,7 +150,6 @@ def indexer_context_scores( Q_TOKEN_STRIDE=idx_q.stride(0), Q_HEAD_STRIDE=idx_q.stride(1), TABLE_STRIDE=block_table.stride(0), - TOKENS=tokens, HEADS=heads, QUERY_LEN=max_query_len, LOCAL_BLOCKS=local, diff --git a/atom/model_ops/qwen4_exp/ops/ple.py b/atom/model_ops/qwen4_exp/ops/ple.py index d3653d63e9..62342f6dcf 100644 --- a/atom/model_ops/qwen4_exp/ops/ple.py +++ b/atom/model_ops/qwen4_exp/ops/ple.py @@ -373,6 +373,7 @@ def ple_gate( return out +# Keep Triton's bounded singleton/alignment variants for the request search. @triton.jit def _conv( X, @@ -392,7 +393,7 @@ def _conv( IS: tl.constexpr, OS: tl.constexpr, SPEC: tl.constexpr, - R: tl.constexpr, + R, B: tl.constexpr, ): token = tl.program_id(0) diff --git a/atom/model_ops/qwen4_exp/ops/qsa.py b/atom/model_ops/qwen4_exp/ops/qsa.py index ce842984f8..9ed62020a0 100644 --- a/atom/model_ops/qwen4_exp/ops/qsa.py +++ b/atom/model_ops/qwen4_exp/ops/qsa.py @@ -85,7 +85,7 @@ def _kernel_config( return 32, 4, 1, 3, sub_group -@triton.jit +@triton.jit(do_not_specialize=["N", "REAL"]) def _draft_decode_metadata( Lengths, Tables, @@ -94,8 +94,8 @@ def _draft_decode_metadata( Positions, Requests, Compressed, - N: tl.constexpr, - REAL: tl.constexpr, + N, + REAL, TS: tl.constexpr, PAGE: tl.constexpr, RATIO: tl.constexpr, diff --git a/atom/model_ops/triton_fused_qkv_quant.py b/atom/model_ops/triton_fused_qkv_quant.py index efbb44f207..d6285b45f3 100644 --- a/atom/model_ops/triton_fused_qkv_quant.py +++ b/atom/model_ops/triton_fused_qkv_quant.py @@ -322,30 +322,36 @@ def fused_qkv_per_tensor_quant(q, k, v, *, k_rope=None): return (*outputs, *(scales[i : i + 1] for i in range(5))) -@triton.jit +@triton.jit(do_not_specialize=["TOKENS"]) def _kv_amax( K, V, Partial, - NK: tl.constexpr, - NV: tl.constexpr, + TOKENS: tl.int64, + K_ROW: tl.constexpr, + V_ROW: tl.constexpr, PARTS: tl.constexpr, BLOCK: tl.constexpr, ): part = tl.program_id(0) kind = tl.program_id(1) X = K if kind == 0 else V - n = NK if kind == 0 else NV - offsets = part * BLOCK + tl.arange(0, BLOCK) + # Fixed row widths preserve vector alignment for every token count. + row = K_ROW if kind == 0 else V_ROW + n = TOKENS * row acc = tl.full((BLOCK,), 0, tl.float32) - for start in range(tl.cdiv(n, PARTS * BLOCK)): - index = offsets + start * PARTS * BLOCK + # The small-input configuration benefits from overlapping two loads; + # larger reductions favor the occupancy of the compact loop. + for start in tl.range( + part * BLOCK, n, PARTS * BLOCK, loop_unroll_factor=2 if PARTS == 256 else 1 + ): + index = tl.multiple_of(start, BLOCK) + tl.arange(0, BLOCK) x = tl.load(X + index, index < n, 0).to(tl.float32) acc = tl.maximum(acc, tl.abs(x)) tl.store(Partial + kind * PARTS + part, tl.max(acc, 0)) -@triton.jit +@triton.jit(do_not_specialize=["TOKENS"]) def _kv_quant( K, V, @@ -353,18 +359,20 @@ def _kv_quant( V8, Partial, Scales, - NK: tl.constexpr, - NV: tl.constexpr, + TOKENS: tl.int64, + K_ROW: tl.constexpr, + V_ROW: tl.constexpr, PARTS: tl.constexpr, BLOCK: tl.constexpr, - WORKERS: tl.constexpr, SINGLE_PASS: tl.constexpr, ): worker = tl.program_id(0) kind = tl.program_id(1) X = K if kind == 0 else V Y = K8 if kind == 0 else V8 - n = NK if kind == 0 else NV + # Fixed row widths preserve vector alignment for every token count. + row = K_ROW if kind == 0 else V_ROW + n = TOKENS * row offsets = worker * BLOCK + tl.arange(0, BLOCK) if SINGLE_PASS: x = tl.load(X + offsets, offsets < n, 0).to(tl.float32) @@ -387,8 +395,8 @@ def _kv_quant( value = tl.minimum(tl.maximum(x * inv, -448.0), 448.0) tl.store(Y + offsets, value, offsets < n) else: - for start in range(tl.cdiv(n, WORKERS * BLOCK)): - index = offsets + start * WORKERS * BLOCK + for start in range(worker * BLOCK, n, tl.num_programs(0) * BLOCK): + index = tl.multiple_of(start, BLOCK) + tl.arange(0, BLOCK) x = tl.load(X + index, index < n, 0).to(tl.float32) value = tl.minimum(tl.maximum(x * inv, -448.0), 448.0) tl.store(Y + index, value, index < n) @@ -432,16 +440,20 @@ def fused_kv_per_tensor_quant(k: torch.Tensor, v: torch.Tensor): if not k.is_contiguous() or not v.is_contiguous(): raise ValueError("K/V must be contiguous") - nk, nv = k.numel(), v.numel() - n = max(nk, nv) + tokens = k.shape[0] + k_row, v_row = k.shape[1] * k.shape[2], v.shape[1] * v.shape[2] + n = tokens * max(k_row, v_row) parts, block, workers, warps = _kv_config(n) k8 = torch.empty(k.shape, dtype=torch.float8_e4m3fn, device=k.device) v8 = torch.empty(v.shape, dtype=torch.float8_e4m3fn, device=v.device) scales = torch.empty(2, dtype=torch.float32, device=k.device) - partial = torch.empty((2, parts), dtype=torch.float32, device=k.device) + partial = None single = n <= 8192 if not single: - _kv_amax[(parts, 2)](k, v, partial, nk, nv, parts, block, num_warps=warps) + partial = torch.empty((2, parts), dtype=torch.float32, device=k.device) + _kv_amax[(parts, 2)]( + k, v, partial, tokens, k_row, v_row, parts, block, num_warps=warps + ) _kv_quant[(workers, 2)]( k, v, @@ -449,11 +461,11 @@ def fused_kv_per_tensor_quant(k: torch.Tensor, v: torch.Tensor): v8, partial, scales, - nk, - nv, + tokens, + k_row, + v_row, parts, block, - workers, single, num_warps=warps, ) diff --git a/atom/model_ops/triton_gemma_rmsnorm.py b/atom/model_ops/triton_gemma_rmsnorm.py index 848c2e087b..6e403bdb16 100644 --- a/atom/model_ops/triton_gemma_rmsnorm.py +++ b/atom/model_ops/triton_gemma_rmsnorm.py @@ -20,7 +20,7 @@ # ── Triton kernel ──────────────────────────────────────────────────────────── -@triton.jit +@triton.jit(do_not_specialize=["n_rows"]) def _gemma_rmsnorm_kernel( input_ptr, output_ptr, @@ -34,7 +34,6 @@ def _gemma_rmsnorm_kernel( epsilon, HAS_RESIDUAL: tl.constexpr, BLOCK_SIZE: tl.constexpr, - NUM_PRGMS: tl.constexpr, GROUPS: tl.constexpr = 1, ): """Fused add + GemmaRMSNorm (weight offset x * (1 + w)). @@ -50,7 +49,7 @@ def _gemma_rmsnorm_kernel( g = tl.load(g_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32) g = g + 1.0 # Gemma offset - for row_idx in tl.range(row_start, n_rows, NUM_PRGMS, num_stages=2): + for row_idx in tl.range(row_start, n_rows, tl.num_programs(0), num_stages=2): if GROUPS > 1: # Each group has its own affine weights in the full-width vector. g = ( @@ -134,7 +133,6 @@ def gemma_rmsnorm_triton(x, weight, eps, residual, group_size=None): eps, HAS_RESIDUAL=has_residual, BLOCK_SIZE=BLOCK_SIZE, - NUM_PRGMS=NUM_PRGMS, GROUPS=groups, ) diff --git a/atom/model_ops/triton_mrope.py b/atom/model_ops/triton_mrope.py index 126d572b1d..a5622177b9 100644 --- a/atom/model_ops/triton_mrope.py +++ b/atom/model_ops/triton_mrope.py @@ -16,7 +16,7 @@ from torch import nn -@triton.jit +@triton.jit(do_not_specialize=["pos_stride_row"]) def _mrope_qk_kernel( q_ptr, k_ptr, @@ -29,7 +29,7 @@ def _mrope_qk_kernel( k_stride_t: tl.constexpr, q_out_stride_t: tl.constexpr, k_out_stride_t: tl.constexpr, - pos_stride_row: tl.constexpr, + pos_stride_row, cos_stride_pos: tl.constexpr, sin_stride_pos: tl.constexpr, num_q_heads: tl.constexpr, @@ -102,7 +102,7 @@ def _mrope_qk_kernel( tl.store(k_out_ptr + k_base_out + d, out, mask=mask & ~is_q) -@triton.jit +@triton.jit(do_not_specialize=["pos_stride_row", "num_tokens"]) def _mrope_qk_tiled_kernel( q_ptr, k_ptr, @@ -115,7 +115,7 @@ def _mrope_qk_tiled_kernel( k_stride_t: tl.constexpr, q_out_stride_t: tl.constexpr, k_out_stride_t: tl.constexpr, - pos_stride_row: tl.constexpr, + pos_stride_row, cos_stride_pos: tl.constexpr, sin_stride_pos: tl.constexpr, num_tokens, diff --git a/tests/model_ops/test_triton_dynamic_shapes.py b/tests/model_ops/test_triton_dynamic_shapes.py new file mode 100644 index 0000000000..69501c803e --- /dev/null +++ b/tests/model_ops/test_triton_dynamic_shapes.py @@ -0,0 +1,361 @@ +# SPDX-License-Identifier: MIT +"""Dynamic batch/length bounds must preserve results without new JIT variants.""" + +from contextlib import contextmanager +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + +pytest.importorskip("triton") +pytest.importorskip("aiter") +pytestmark = pytest.mark.skipif( + not torch.version.hip or not torch.cuda.is_available(), reason="requires a ROCm GPU" +) + + +@contextmanager +def one_variant(*kernels): + # Isolate this assertion from shapes warmed by other tests. Restore their + # entries afterward; no disk cache or compiled module is removed. + caches = [k.device_caches[torch.cuda.current_device()][0] for k in kernels] + saved = [dict(cache) for cache in caches] + for cache in caches: + cache.clear() + try: + yield + torch.cuda.synchronize() + for kernel, cache in zip(kernels, caches): + assert len(cache) == 1, (kernel.__name__, len(cache)) + finally: + for cache, entries in zip(caches, saved): + cache.clear() + cache.update(entries) + + +@pytest.mark.parametrize("groups", [1, 4]) +def test_gemma_rows_reuse_kernel(groups): + from atom.model_ops import triton_gemma_rmsnorm as m + + torch.manual_seed(41) + weight = torch.randn(groups * 128, device="cuda", dtype=torch.bfloat16) + with one_variant(m._gemma_rmsnorm_kernel): + for rows in (1, 15, 16, 17, 127, 305): + x = torch.randn(rows, groups * 128, device="cuda", dtype=torch.bfloat16) + residual = torch.randn_like(x) + out, residual_out = m.gemma_rmsnorm_triton(x, weight, 1e-6, residual, 128) + combined = (x.float() + residual.float()).reshape(rows, groups, 128) + expected = combined * torch.rsqrt( + combined.square().mean(-1, keepdim=True) + 1e-6 + ) + expected *= 1 + weight.float().view(groups, 128) + torch.testing.assert_close(out, expected.reshape_as(x).to(x.dtype)) + torch.testing.assert_close(residual_out, combined.reshape_as(x).to(x.dtype)) + + +@pytest.mark.parametrize("tiled", [False, True]) +def test_mrope_position_stride_reuses_kernel(tiled): + from atom.model_ops import triton_mrope as m + + torch.manual_seed(42) + angles = torch.randn(64, 32, device="cuda") + rotary = SimpleNamespace( + mrope_section=[8, 12, 12], + mrope_interleaved=True, + rotary_dim=64, + is_neox_style=True, + cos_cache=angles.cos().to(torch.bfloat16), + sin_cache=angles.sin().to(torch.bfloat16), + ) + columns = torch.arange(32, device="cuda") + axes = torch.where(columns % 3 == 1, 1, torch.where(columns % 3 == 2, 2, 0)) + kernel = m._mrope_qk_tiled_kernel if tiled else m._mrope_qk_kernel + with one_variant(kernel): + for n in ((128, 129, 255, 256, 257) if tiled else (1, 15, 16, 17, 127)): + q = torch.randn(n, 512, device="cuda", dtype=torch.bfloat16) + k = torch.randn(n, 256, device="cuda", dtype=torch.bfloat16) + positions = torch.randint(0, 64, (3, n), device="cuda") + output = m.try_mrope_qk_fused(rotary, positions, q, k, 2, 1, 256) + selected = positions[axes].T + cos = rotary.cos_cache[selected, columns].float()[:, None, :] + sin = rotary.sin_cache[selected, columns].float()[:, None, :] + for x, actual in zip((q, k), output): + heads = x.reshape(n, -1, 256).float() + left, right = heads[..., :32], heads[..., 32:64] + expected = torch.cat( + ( + left * cos - right * sin, + right * cos + left * sin, + heads[..., 64:], + ), + -1, + ) + torch.testing.assert_close( + actual, expected.reshape_as(x).to(x.dtype), atol=0.016, rtol=0.016 + ) + + +@pytest.mark.parametrize( + "sizes,heads,single", + [ + ((17, 18, 19, 20, 21, 22), 1, True), + ((257, 259, 261, 263, 265, 267), 4, False), + ((1025, 2049, 3073, 4097, 6145), 12, False), + ], +) +def test_cached_kv_lengths_reuse_kernels(sizes, heads, single): + from atom.model_ops import triton_fused_qkv_quant as m + + kernels = (m._kv_quant,) if single else (m._kv_amax, m._kv_quant) + with one_variant(*kernels): + for n in sizes: + inputs = [] + for dim, value in ((4, 2), (3, -9)) if single else ((192, 2), (128, -9)): + storage = torch.full( + (n * heads * dim + 128,), 128, device="cuda", dtype=torch.bfloat16 + ) + x = storage[: n * heads * dim].reshape(n, heads, dim) + x.fill_(value) + # Put the maximum in the tail, including odd unrolled trips. + x[-1].mul_(2) + inputs.append(x) + k8, v8, ks, vs = m.fused_kv_per_tensor_quant(*inputs) + for actual, sign in ((k8, 1), (v8, -1)): + assert torch.all(actual[:-1].float() == sign * 224) + assert torch.all(actual[-1].float() == sign * 448) + assert ks.item() == pytest.approx(4 / 448) + assert vs.item() == pytest.approx(18 / 448) + + +@pytest.mark.parametrize("prefill", [False, True]) +def test_mla_dynamic_bounds_preserve_request_regions(prefill): + from atom.model_ops import attention_mla as m + + kernel = ( + m._convert_req_index_to_global_index_dsa_prefill_kernel + if prefill + else m._convert_req_index_to_global_index_kernel + ) + workspace = torch.empty(64 * 128 + 16, device="cuda", dtype=torch.int32) + dense = torch.arange(64 * 256, device="cuda", dtype=torch.int32) + with one_variant(kernel): + for n in (1, 15, 16, 17, 31, 32, 33): + ptr = torch.arange(n + 1, device="cuda", dtype=torch.int32) + topk = torch.arange(128, device="cuda", dtype=torch.int32).repeat(n, 1) + topk[:, 3] = -1 + topk[:, 5] = 256 + workspace.fill_(7) + if prefill: + table = ( + torch.arange(n * 32, device="cuda", dtype=torch.int32).reshape( + n, 32 + ) + + 100 + ) + out = m.triton_convert_req_index_to_global_index_dsa_prefill( + ptr, + ptr * 125, + ptr[:-1].contiguous(), + topk, + table, + ptr * 256, + PAGE_SIZE=16, + NUM_TOPK_TOKENS=128, + BLOCK_N=128, + out=workspace, + seq_local=True, + ) + base = (ptr[:-1] * 32 + 100) * 16 + else: + out = m.triton_convert_req_index_to_global_index( + ptr, + ptr * 256, + ptr * 125, + dense, + topk, + NUM_TOPK_TOKENS=128, + out=workspace, + ) + base = ptr[:-1] * 256 + expected = base[:, None] + torch.arange( + 125, device="cuda", dtype=torch.int32 + ) + expected[:, 3] = expected[:, 5] = 0 + assert torch.equal(out[: n * 125], expected.flatten()) + assert torch.all(workspace[n * 125 :] == 7) + + +@pytest.mark.parametrize("sizes", [(1,), (15, 17, 31, 33), (16, 32, 48)]) +def test_ple_request_count_reuses_kernel(sizes): + from atom.model_ops.qwen4_exp.ops import ple as m + + torch.manual_seed(43) + weight = torch.randn(32, 3, device="cuda", dtype=torch.bfloat16) + with one_variant(m._conv): + for n in sizes: + x = torch.randn(n * 2, 32, device="cuda", dtype=torch.bfloat16) + slots = torch.arange(n, device="cuda", dtype=torch.int32) + state = torch.zeros(64, 32, 2, device="cuda", dtype=x.dtype) + starts = torch.arange(n + 1, device="cuda", dtype=torch.int32) * 2 + actual = m.dilated_causal_conv1d( + x, + weight, + state, + starts, + slots, + slots, + torch.zeros(n, device="cuda", dtype=torch.bool), + 1, + ) + ref = F.conv1d( + x.reshape(n, 2, 32).transpose(1, 2).float(), + weight[:, None, :].float(), + groups=32, + padding=2, + )[..., :2].to(x.dtype) + ref = F.silu(ref.float()).to(x.dtype).transpose(1, 2).reshape_as(x) + torch.testing.assert_close(actual, ref, atol=0.016, rtol=0.016) + + +def test_qsa_real_requests_change_within_and_between_buckets(): + from atom.model_ops.qwen4_exp.ops import qsa as m + + with one_variant(m._draft_decode_metadata): + for n, real in ((64, 1), (64, 15), (64, 16), (65, 17), (128, 31), (129, 33)): + lengths = torch.full((n,), 19, device="cuda", dtype=torch.int32) + table = torch.arange(n * 32, device="cuda", dtype=torch.int32).reshape( + n, 32 + ) + rejects = torch.full((real,), 3, device="cuda", dtype=torch.int32) + slots, positions, requests, compressed = [ + torch.empty(n, device="cuda", dtype=torch.int64) for _ in range(4) + ] + m.qsa_draft_decode_metadata( + lengths, + table, + rejects, + slots, + positions, + requests, + compressed, + real, + 16, + 4, + ) + ids = torch.arange(real, device="cuda") + assert torch.equal(slots[:real], ids * 512 + 15) + assert torch.equal(compressed[:real], (ids * 512 + 15) // 4) + assert torch.equal(requests[:real], ids) + assert torch.all(positions[:real] == 15) and torch.all(lengths[:real] == 16) + for x in (slots, positions, requests, compressed): + assert torch.all(x[real:] == -1) + assert not lengths[real:].count_nonzero() + + +def test_m3_context_and_batch_reuse_kernels(): + from atom.model_ops.minimax_m3 import indexer_candidate_exchange as ex + from atom.model_ops.minimax_m3 import indexer_context_parallel as cp + + cache = torch.zeros(64, 128, 128, device="cuda", dtype=torch.bfloat16) + with one_variant(cp._context_score, ex._local_topk, ex._merge_topk): + for n, blocks in ((1, 33), (15, 35), (16, 37), (17, 39), (31, 41)): + q = torch.ones(n, 4, 128, device="cuda", dtype=torch.bfloat16) + table = torch.arange(64, device="cuda", dtype=torch.int32).repeat(n, 1) + lengths = torch.full((n,), blocks * 128, device="cuda", dtype=torch.int32) + scores = cp.indexer_context_scores( + q, cache, table, lengths, blocks * 128, 0, 4, 1, 0.1 + ) + assert not scores.count_nonzero() + keys = ex.local_candidate_keys(scores, lengths, 4, 0, 4, 1, 64) + indices, _, _ = ex.merge_candidate_keys(keys, table, lengths, 4, 0, 0, 1) + assert torch.all((indices >= 0) & (indices < blocks)) + + +def test_packed_gather_bounds_reuse_kernel(): + from atom.model_ops.attentions.deepseek_v41 import packed_rows as m + from atom.model_ops.blockscale import quantize_fp8 + + torch.manual_seed(44) + x = torch.randn(128, 32, device="cuda", dtype=torch.bfloat16) + pool = m.pack_rows(*quantize_fp8(x)).flatten() + ref = quantize_fp8(x, dequantize=True) + order = torch.randperm(128, device="cuda") + addresses = ((order * 33 * 2) | 1).contiguous() + addresses[111:] = 2**60 + ptr = torch.tensor([0, 7, 23, 58, 111], device="cuda", dtype=torch.int32) + with one_variant(m._gather_prefix): + for capacity, first, end in ( + (1, 0, 1), + (15, 1, 2), + (16, 1, 2), + (17, 2, 4), + (33, 0, 4), + (128, 0, 4), + ): + backing = torch.full((capacity + 16, 32), 7, device="cuda", dtype=x.dtype) + out = backing[:capacity] + m.gather_prefix_rows(pool, addresses, ptr, out, first, end) + start, stop = (0, 7, 23, 58, 111)[first], (0, 7, 23, 58, 111)[end] + count = min(capacity, stop - start) + assert torch.equal(out[:count], ref[order[start : start + count]]) + assert not out[count:].count_nonzero() + assert torch.all(backing[capacity:] == 7) + + +def test_cached_kv_large_lengths_keep_a_single_integer_signature(): + from atom.model_ops import triton_fused_qkv_quant as m + + # Compile only: exercise signed/unsigned boundaries without making the + # regular suite allocate multi-GB tensors or launching into tiny storage. + x = torch.empty(1, device="cuda", dtype=torch.bfloat16) + y = torch.empty(1, device="cuda", dtype=torch.float8_e4m3fn) + partial = torch.empty((2, 512), device="cuda") + scales = torch.empty(2, device="cuda") + with one_variant(m._kv_amax, m._kv_quant): + for tokens in (2**31 - 1, 2**31, 2**32 + 1): + m._kv_amax.warmup( + x, x, partial, tokens, 192, 128, 512, 8192, grid=(512, 2), num_warps=8 + ) + m._kv_quant.warmup( + x, + x, + y, + y, + partial, + scales, + tokens, + 192, + 128, + 512, + 8192, + False, + grid=(4096, 2), + num_warps=8, + ) + + +@pytest.mark.parametrize("alignment", [0, 1]) +def test_m3_local_topk_reuses_each_alignment_class(alignment): + from atom.model_ops.minimax_m3 import indexer_candidate_exchange as m + + lengths = torch.full((4,), 1048576, device="cuda", dtype=torch.int32) + with one_variant(m._local_topk): + for blocks in (1040, 1056, 1072): + scores = torch.zeros((4, 4, blocks + alignment), device="cuda") + m.local_candidate_keys(scores, lengths, 16, 0, 4, 1, 8192) + + +@pytest.mark.parametrize("blocks", [1025, 1536, 2048, 2049]) +def test_m3_local_topk_long_rows_match_torch(blocks): + from atom.model_ops.minimax_m3 import indexer_candidate_exchange as m + + torch.manual_seed(45) + # Unique signed scores make the expected ordering independent of tie rules. + scores = torch.randperm(16 * blocks, device="cuda").reshape(4, 4, blocks).float() + scores -= 8 * blocks + lengths = torch.full((4,), blocks * 4 * 128, device="cuda", dtype=torch.int32) + keys = m.local_candidate_keys(scores, lengths, 16, 0, 4, 1, blocks * 4) + indices = (keys & 0xFFFF) - 1 + expected = scores.topk(16, dim=-1).indices * 4 + assert torch.equal(indices, expected) From f94eb2b82487e085c9e2071f4a963eac691e0281 Mon Sep 17 00:00:00 2001 From: Lingpeng Jin <103567126+valarLip@users.noreply.github.com> Date: Fri, 25 Sep 2026 05:55:32 +0000 Subject: [PATCH 5/5] fix(distributed): own rendezvous store before spawning workers --- atom/config.py | 7 +- atom/model_engine/engine_core_mgr.py | 43 ++++- atom/model_engine/model_runner.py | 31 +++- docs/configuration_guide.md | 2 +- docs/distributed_guide.md | 11 +- tests/test_distributed_store.py | 251 +++++++++++++++++++++++++++ 6 files changed, 327 insertions(+), 18 deletions(-) create mode 100644 tests/test_distributed_store.py diff --git a/atom/config.py b/atom/config.py index e1b6d87333..73933282ae 100644 --- a/atom/config.py +++ b/atom/config.py @@ -26,7 +26,7 @@ LayerQuantConfig, get_quant_parser, ) -from atom.utils import envs, get_open_port +from atom.utils import envs from atom.utils.distributed.utils import stateless_init_torch_distributed_process_group if TYPE_CHECKING: @@ -1018,7 +1018,10 @@ class ParallelConfig: data_parallel_master_port: int = 29500 """Port of the data parallel master.""" - data_parallel_base_port: int = get_open_port() + data_parallel_base_port: int = 0 + """Model-runner rendezvous port. Zero requests an OS-assigned port locally.""" + _managed_distributed_store: bool = field(default=False, init=False, repr=False) + """CoreManager owns the store; every model runner connects as a client.""" data_parallel_master_ip: str = "127.0.0.1" diff --git a/atom/model_engine/engine_core_mgr.py b/atom/model_engine/engine_core_mgr.py index 766cf673f6..68f07ecda3 100644 --- a/atom/model_engine/engine_core_mgr.py +++ b/atom/model_engine/engine_core_mgr.py @@ -196,6 +196,7 @@ def _init_shared_state( # is only safe to admit if it fits wherever the router sends it. self.max_pool_tokens: int | None = None self.engine_core_processes = [] + self._distributed_stores = [] self.input_sockets = [] self.output_sockets = [] self.engine_core_identities = [] @@ -264,6 +265,35 @@ def _init_shared_state( # scraping costs no round trip and cannot time out. self.latest_metrics: dict[int, dict] = {} + def _start_distributed_store( + self, config: Config, *, multinode: bool = False + ) -> None: + """Bind the actual rendezvous server before publishing its port to workers.""" + from torch.distributed import TCPStore + + pc = config.parallel_config + if multinode and pc.data_parallel_base_port == 0: + raise ValueError( + "Multi-node DP requires the same nonzero --data-parallel-base-port " + "(or ATOM_DP_BASE_PORT) on every node." + ) + if not multinode or pc.data_parallel_rank == 0: + store = TCPStore( + pc.data_parallel_master_ip, + pc.data_parallel_base_port, + is_master=True, + wait_for_workers=False, + ) + self._distributed_stores.append(store) + pc.data_parallel_base_port = store.port + logger.info( + "%s: model-runner TCPStore listening on %s:%d", + self.label, + pc.data_parallel_master_ip, + store.port, + ) + pc._managed_distributed_store = True + def __init__(self, config: Config): pp_size = config.pipeline_parallel_size self.pp_size = pp_size @@ -344,6 +374,7 @@ def __init__(self, config: Config): local_dp_ranks = [] try: + self._start_distributed_store(config, multinode=multinode) for engine_index in range(self.local_engine_count): assignment_index = engine_index // self.pp_size dp_rank, local_dp_rank = rank_assignments[assignment_index] @@ -803,6 +834,8 @@ def close(self): except (ValueError, OSError): pass + # Release the rendezvous server after stopping the local engine processes. + self._distributed_stores.clear() logger.info(f"{self.label}: All EngineCores shut down") def _send_request(self, dp_rank: int, payload: bytes) -> None: @@ -1517,8 +1550,6 @@ def __init__(self, config: Config): self._cu_shm = None # Build per-process configs. - from atom.utils import get_open_port as _get_open_port - prefill_config = copy.deepcopy(config) if config.disagg_prefill_max_num_seqs is not None: prefill_config.max_num_seqs = config.disagg_prefill_max_num_seqs @@ -1529,10 +1560,8 @@ def __init__(self, config: Config): prefill_config.disagg_weight_ack_addr = weight_ack_addr prefill_config.disagg_kvcache_ipc_addr = kvcache_ipc_addr prefill_config.disagg_cu_shm_name = cu_shm_name - # Give prefill a distinct distributed rendezvous port so it doesn't - # collide with decode's data_parallel_base_port (both deep-copy the - # same port from config). - prefill_config.parallel_config.data_parallel_base_port = _get_open_port() + # Prefill gets its own server; decode honors the configured port. + prefill_config.parallel_config.data_parallel_base_port = 0 decode_config = copy.deepcopy(config) decode_config.disagg_d2p_addr = d2p_addr @@ -1613,6 +1642,8 @@ def _connect_proc(proc, in_addr, out_addr, ctrl_addr, name): logger.info(f"{self.label}: {name} process started and connected") try: + self._start_distributed_store(decode_config) + self._start_distributed_store(prefill_config) # Start both processes simultaneously. Prefill binds the bootstrap # PUSH socket and blocks on send() until decode connects and calls # recv() — they rendezvous naturally without any sequential ordering. diff --git a/atom/model_engine/model_runner.py b/atom/model_engine/model_runner.py index 6f4ceae01a..4129209400 100644 --- a/atom/model_engine/model_runner.py +++ b/atom/model_engine/model_runner.py @@ -932,6 +932,9 @@ def _setup_device_and_distributed(self, rank: int, config: Config): dp_rank_local = config.parallel_config.data_parallel_rank_local or 0 pp_rank = config.parallel_config.pipeline_parallel_rank pp_size = config.pipeline_parallel_size + if pp_size > 1: + # Reject before any collective can wait for nonexistent ranks. + reject_simulated_tp(config, "pipeline parallel") # tp_world_size: how many GPUs this stage actually occupies. stage_span = config.tp_world_size * config.prefill_context_parallel_size engine_index = dp_rank_local * pp_size + pp_rank @@ -956,16 +959,30 @@ def _setup_device_and_distributed(self, rank: int, config: Config): config.parallel_config.data_parallel_master_ip, config.parallel_config.data_parallel_base_port, ) - # Both branches handle simulated TP: the PP path only to reject it, - # since it would otherwise deadlock on a group sized for absent ranks. + dp_size = config.parallel_config.data_parallel_size + world_size = dp_size * pp_size * stage_span + dp_rank = config.parallel_config.data_parallel_rank + global_rank = (dp_rank * pp_size + pp_rank) * stage_span + rank + if ( + config.parallel_config._managed_distributed_store + and not torch.distributed.is_initialized() + ): + # Preserve AITER's environment setup when preinitializing its group. + os.environ.setdefault( + "HIP_VISIBLE_DEVICES", ",".join(map(str, range(world_size))) + ) + store = torch.distributed.TCPStore( + config.parallel_config.data_parallel_master_ip, + config.parallel_config.data_parallel_base_port, + is_master=False, + ) + torch.distributed.init_process_group( + backend="nccl", store=store, rank=global_rank, world_size=world_size + ) + # AITER reuses the default group and creates the model-parallel groups. if config.pipeline_parallel_size > 1: from atom.distributed.pp_comm import init_pp_aware_dist_env - reject_simulated_tp(config, "pipeline parallel") - dp_size = config.parallel_config.data_parallel_size - world_size = dp_size * pp_size * stage_span - dp_rank = config.parallel_config.data_parallel_rank - global_rank = (dp_rank * pp_size + pp_rank) * stage_span + rank # No local_rank here, unlike the non-PP branch below. Safe only # because PP is single-node today: CoreManager rejects multi-node # DP when pp_size > 1, and asserts PP+DP out entirely, so diff --git a/docs/configuration_guide.md b/docs/configuration_guide.md index 569e248b4c..3ea1d0095c 100644 --- a/docs/configuration_guide.md +++ b/docs/configuration_guide.md @@ -263,7 +263,7 @@ Defined in `atom/config.py`. Controls data parallelism. Environment variables | `data_parallel_rank` | `int` | `0` | First **global** DP rank owned by this node; overridden by `ATOM_DP_RANK` | | `data_parallel_rank_local` | `Optional[int]` | `None` | Local rank within the data-parallel group (SPMD mode); overridden by `ATOM_DP_RANK_LOCAL` | | `data_parallel_master_port` | `int` | `29500` | Port used by the data-parallel master for process group initialization | -| `data_parallel_base_port` | `int` | `get_open_port()` | Base port for data-parallel communication (dynamically assigned) | +| `data_parallel_base_port` | `int` | `0` | Model-runner TCPStore port; automatically bound for single-node runs, explicitly shared across nodes | | `data_parallel_master_ip` | `str` | `"127.0.0.1"` | IP address of the data-parallel master | **Computed property:** diff --git a/docs/distributed_guide.md b/docs/distributed_guide.md index 9d366f7ae8..23088368c6 100644 --- a/docs/distributed_guide.md +++ b/docs/distributed_guide.md @@ -534,18 +534,25 @@ There is no enable flag. Each has an `ATOM_DP_*` environment equivalent (`ATOM_DP_SIZE_LOCAL`, `ATOM_DP_RANK`, …) for launch scripts; explicit flags win. +Set the same nonzero `--data-parallel-base-port` on every node. The coordinator +binds and holds the model-runner TCPStore before spawning workers; all model +runners connect as clients. On a single node, the default `0` lets the OS +allocate the port when the store starts, so there is no probe-and-release gap. + ### Example — 2 nodes, 4 DP ranks each, TP2 ```bash # Node 0 (coordinator, 10.0.0.1) — serves the API python -m atom.entrypoints.openai_server --model -tp 2 \ --data-parallel-size 8 --data-parallel-size-local 4 --data-parallel-rank 0 \ - --data-parallel-master-ip 10.0.0.1 --data-parallel-master-port 29500 + --data-parallel-master-ip 10.0.0.1 --data-parallel-master-port 29500 \ + --data-parallel-base-port 29501 # Node 1 — engines only, no API server python -m atom.entrypoints.openai_server --model -tp 2 \ --data-parallel-size 8 --data-parallel-size-local 4 --data-parallel-rank 4 \ - --data-parallel-master-ip 10.0.0.1 --data-parallel-master-port 29500 + --data-parallel-master-ip 10.0.0.1 --data-parallel-master-port 29500 \ + --data-parallel-base-port 29501 ``` Node 0 owns global DP ranks 0-3, node 1 owns 4-7. Each node's *local* ranks diff --git a/tests/test_distributed_store.py b/tests/test_distributed_store.py new file mode 100644 index 0000000000..f88e0df287 --- /dev/null +++ b/tests/test_distributed_store.py @@ -0,0 +1,251 @@ +# SPDX-License-Identifier: MIT +"""Rendezvous ports stay owned from allocation through worker shutdown.""" + +import copy +import errno +import multiprocessing +import socket +import time +from datetime import timedelta +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist + +from atom.config import ParallelConfig +from atom.model_engine.engine_core_mgr import CoreManager + + +@pytest.fixture +def manager(monkeypatch): + for name in ( + "ATOM_DP_SIZE", + "ATOM_DP_SIZE_LOCAL", + "ATOM_DP_RANK", + "ATOM_DP_RANK_LOCAL", + "ATOM_DP_MASTER_IP", + "ATOM_DP_MASTER_PORT", + "ATOM_DP_BASE_PORT", + ): + monkeypatch.delenv(name, raising=False) + manager = CoreManager.__new__(CoreManager) + manager._init_shared_state( + SimpleNamespace(dp_load_balance="round_robin"), + label="test", + local_engine_count=0, + ) + try: + yield manager + finally: + manager.close() + manager.ctx.term() + + +def _assert_port_owned(port): + with socket.socket() as contender: + with pytest.raises(OSError) as error: + contender.bind(("127.0.0.1", port)) + assert error.value.errno == errno.EADDRINUSE + + +def _spawn_workers(target, args_per_rank): + context = multiprocessing.get_context("spawn") + processes = [context.Process(target=target, args=args) for args in args_per_rank] + try: + for process in processes: + process.start() + deadline = time.monotonic() + 120 + for process in processes: + process.join(timeout=max(0, deadline - time.monotonic())) + assert [p.exitcode for p in processes] == [0] * len(processes) + finally: + for process in processes: + if process.is_alive(): + process.kill() + if process.pid is not None: + process.join(timeout=5) + process.close() + + +def _gloo_worker(pc, rank, world_size): + store = dist.TCPStore( + pc.data_parallel_master_ip, + pc.data_parallel_base_port, + is_master=False, + timeout=timedelta(seconds=30), + ) + dist.init_process_group( + "gloo", + store=store, + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=30), + ) + try: + value = torch.tensor(rank + 1) + dist.all_reduce(value) + assert value.item() == world_size * (world_size + 1) // 2 + finally: + dist.destroy_process_group() + + +def test_store_survives_workers_and_is_released_on_close(manager): + config = SimpleNamespace(parallel_config=ParallelConfig()) + assert config.parallel_config.data_parallel_base_port == 0 + manager._start_distributed_store(config) + pc = copy.deepcopy(config.parallel_config) + assert pc._managed_distributed_store + _assert_port_owned(pc.data_parallel_base_port) + _spawn_workers(_gloo_worker, [(pc, rank, 2) for rank in range(2)]) + # Rank zero exiting must not destroy the coordinator's listener. + _assert_port_owned(pc.data_parallel_base_port) + manager.close() + replacement = dist.TCPStore( + pc.data_parallel_master_ip, + pc.data_parallel_base_port, + is_master=True, + wait_for_workers=False, + ) + assert replacement.port == pc.data_parallel_base_port + + +def test_separate_stores_have_distinct_ports_and_keys(manager): + configs = [SimpleNamespace(parallel_config=ParallelConfig()) for _ in range(2)] + for config in configs: + manager._start_distributed_store(config) + ports = [config.parallel_config.data_parallel_base_port for config in configs] + assert ports[0] != ports[1] + clients = [dist.TCPStore("127.0.0.1", port, is_master=False) for port in ports] + for i, client in enumerate(clients): + client.set("same-key", str(i)) + assert [client.get("same-key") for client in clients] == [b"0", b"1"] + + +def test_fixed_port_conflict_fails_without_changing_port(manager): + owner = dist.TCPStore("127.0.0.1", 0, is_master=True, wait_for_workers=False) + pc = ParallelConfig(data_parallel_base_port=owner.port) + with pytest.raises(dist.DistNetworkError, match="EADDRINUSE"): + manager._start_distributed_store(SimpleNamespace(parallel_config=pc)) + assert pc.data_parallel_base_port == owner.port + assert not pc._managed_distributed_store + assert not manager._distributed_stores + + +def test_remote_node_uses_coordinator_store(manager): + owner = dist.TCPStore("127.0.0.1", 0, is_master=True, wait_for_workers=False) + pc = ParallelConfig( + data_parallel_size=2, + data_parallel_size_local=1, + data_parallel_rank=1, + data_parallel_base_port=owner.port, + ) + manager._start_distributed_store( + SimpleNamespace(parallel_config=pc), multinode=True + ) + assert pc._managed_distributed_store + assert not manager._distributed_stores + client = dist.TCPStore("127.0.0.1", pc.data_parallel_base_port, is_master=False) + client.set("remote", "connected") + assert owner.get("remote") == b"connected" + + +@pytest.mark.parametrize("rank", [0, 1]) +def test_multinode_requires_shared_fixed_port(manager, rank): + pc = ParallelConfig( + data_parallel_size=2, + data_parallel_size_local=1, + data_parallel_rank=rank, + ) + with pytest.raises(ValueError, match="same nonzero --data-parallel-base-port"): + manager._start_distributed_store( + SimpleNamespace(parallel_config=pc), multinode=True + ) + assert not manager._distributed_stores + + +def _model_runner_worker(pc, layout, global_rank): + from aiter import destroy_dist_env + from aiter.dist.parallel_state import ( + get_dp_group, + get_pcp_group, + get_pp_group, + get_tp_group, + ) + + from atom.model_engine.model_runner import ModelRunner + + tp, pp, dp, pcp, dcp = layout + stage_span = tp * pcp + pc.data_parallel_rank = global_rank // (pp * stage_span) + pc.data_parallel_rank_local = pc.data_parallel_rank + pc.pipeline_parallel_rank = global_rank // stage_span % pp + config = SimpleNamespace( + parallel_config=pc, + tensor_parallel_size=tp, + tp_world_size=tp, + pipeline_parallel_size=pp, + prefill_context_parallel_size=pcp, + decode_context_parallel_size=dcp, + master_addr="127.0.0.1", + port=0, + ) + runner = ModelRunner.__new__(ModelRunner) + runner.config = config + runner._setup_device_and_distributed(global_rank % stage_span, config) + try: + assert dist.get_rank() == global_rank + assert dist.get_world_size() == tp * pp * dp * pcp + assert get_tp_group().world_size == tp + assert get_pp_group().world_size == pp + assert get_dp_group().world_size == dp + assert get_pcp_group().world_size == pcp + value = torch.tensor(global_rank + 1, device=runner.device) + dist.all_reduce(value) + world_size = dist.get_world_size() + assert value.item() == world_size * (world_size + 1) // 2 + torch.cuda.synchronize() + finally: + destroy_dist_env() + + +@pytest.mark.parametrize( + "layout", + [ + (2, 1, 1, 1, 1), + (1, 1, 2, 1, 1), + (2, 1, 2, 1, 1), + (1, 2, 1, 1, 1), + (1, 1, 1, 2, 1), + (2, 1, 1, 1, 2), + (1, 1, 8, 1, 1), + ], + ids=["tp", "dp", "tp_dp", "pp", "pcp", "dcp", "dp8"], +) +def test_model_runner_connects_to_managed_store(manager, layout): + tp, pp, dp, pcp, _ = layout + world_size = tp * pp * dp * pcp + if not torch.version.hip or torch.cuda.device_count() < world_size: + pytest.skip(f"requires {world_size} ROCm GPUs and AITER") + pytest.importorskip("aiter") + config = SimpleNamespace(parallel_config=ParallelConfig(data_parallel_size=dp)) + manager._start_distributed_store(config) + _spawn_workers( + _model_runner_worker, + [(config.parallel_config, layout, rank) for rank in range(world_size)], + ) + _assert_port_owned(config.parallel_config.data_parallel_base_port) + + +def test_standalone_model_runner_keeps_rank_zero_rendezvous(manager): + if not torch.version.hip or torch.cuda.device_count() < 2: + pytest.skip("requires 2 ROCm GPUs and AITER") + pytest.importorskip("aiter") + from atom.utils import get_open_port + + # Standalone callers choose a shared endpoint without a CoreManager. + pc = ParallelConfig(data_parallel_base_port=get_open_port()) + assert not pc._managed_distributed_store + _spawn_workers( + _model_runner_worker, [(pc, (2, 1, 1, 1, 1), rank) for rank in range(2)] + )