diff --git a/examples/hstu/configs/__init__.py b/examples/hstu/configs/__init__.py index f6cc11d95..4747e656f 100644 --- a/examples/hstu/configs/__init__.py +++ b/examples/hstu/configs/__init__.py @@ -1,4 +1,4 @@ -from . import hstu_config, inference_config, task_config +from . import hstu_config, inference_config from .hstu_config import ( HSTUConfig, HSTULayerType, @@ -13,7 +13,21 @@ InferenceHSTUConfig, get_inference_hstu_config, ) -from .task_config import RankingConfig, RetrievalConfig + + +def __getattr__(name): + """Load task schemas on demand without importing training dependencies.""" + # Inference configuration and tensor-only layer tests do not require the + # training embedding stack. Load task/embedding schemas only when requested. + if name in ("task_config", "RankingConfig", "RetrievalConfig"): + from importlib import import_module + + module = import_module(".task_config", __name__) + value = module if name == "task_config" else getattr(module, name) + globals()[name] = value + return value + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + __all__ = [ "hstu_config", diff --git a/examples/hstu/configs/inference_config.py b/examples/hstu/configs/inference_config.py index d8612a2fb..c047e192a 100755 --- a/examples/hstu/configs/inference_config.py +++ b/examples/hstu/configs/inference_config.py @@ -81,6 +81,8 @@ class InferenceHSTUConfig: hstu_preprocessing_config (HSTUPreprocessingConfig, optional): HSTU preprocessing config. contextual_max_seqlen (int): The (maximum) length of contextual features. embedding_backend (EmbeddingBackend, optional): Embedding backend to use. + backbone (str): Dense inference backbone, either hstu or transformer. + transformer_ffn_dim (int, optional): Transformer feed-forward width; defaults to four times hidden_size. """ hidden_size: int @@ -103,10 +105,28 @@ class InferenceHSTUConfig: scaling_seqlen: int = -1 embedding_backend: Optional[EmbeddingBackend] = None export_mode: bool = False + # Keep the existing config/API name for compatibility with HSTU callers. + backbone: str = "hstu" + transformer_ffn_dim: Optional[int] = None def __post_init__(self): assert self.is_causal assert self.target_group_size == 1 + if self.backbone not in ("hstu", "transformer"): + raise ValueError(f"Unknown inference backbone: {self.backbone}") + if self.transformer_ffn_dim is not None and self.transformer_ffn_dim <= 0: + raise ValueError("transformer_ffn_dim must be positive") + if self.backbone == "transformer": + for name in ( + "hidden_size", + "num_heads", + "head_dim", + "num_layers", + "max_batch_size", + "max_seq_len", + ): + if getattr(self, name) <= 0: + raise ValueError(f"{name} must be positive") def get_inference_hstu_config( @@ -127,6 +147,9 @@ def get_inference_hstu_config( scaling_seqlen: int = -1, embedding_backend=None, export_mode: Optional[bool] = None, + backbone: str = "hstu", + transformer_ffn_dim: Optional[int] = None, + hstu_preprocessing_config: Optional[HSTUPreprocessingConfig] = None, ) -> InferenceHSTUConfig: """ Create the HSTU configuration. @@ -177,4 +200,7 @@ def get_inference_hstu_config( scaling_seqlen=scaling_seqlen, embedding_backend=embedding_backend, export_mode=export_mode, + backbone=backbone, + transformer_ffn_dim=transformer_ffn_dim, + hstu_preprocessing_config=hstu_preprocessing_config, ) diff --git a/examples/hstu/inference/README.md b/examples/hstu/inference/README.md index a93e6d71a..195fc4618 100644 --- a/examples/hstu/inference/README.md +++ b/examples/hstu/inference/README.md @@ -1,5 +1,9 @@ # HSTU Inference +The ranking workflow also supports a [Transformer dense backbone](transformer.md) +with the same recommendation inputs, embeddings, KV-cache manager and export +entry points. HSTU remains the default. + ## Key Features 1. KV Cache Manager diff --git a/examples/hstu/inference/configs/kuairand_1k_transformer_ranking.gin b/examples/hstu/inference/configs/kuairand_1k_transformer_ranking.gin new file mode 100644 index 000000000..49b32d974 --- /dev/null +++ b/examples/hstu/inference/configs/kuairand_1k_transformer_ranking.gin @@ -0,0 +1,6 @@ +include 'inference/configs/kuairand_1k_inference_ranking.gin' + +# Use a checkpoint trained/adapted to the architecture documented in +# inference/transformer.md. An HSTU checkpoint is not interchangeable. +NetworkArgs.backbone = 'transformer' +NetworkArgs.transformer_ffn_dim = 2048 diff --git a/examples/hstu/inference/inference_gr_ranking.py b/examples/hstu/inference/inference_gr_ranking.py index 27c2b8f5e..aa7ab5b43 100644 --- a/examples/hstu/inference/inference_gr_ranking.py +++ b/examples/hstu/inference/inference_gr_ranking.py @@ -159,6 +159,8 @@ def get_inference_hstu_model( dtype=inference_dtype, position_encoding_config=position_encoding_config, contextual_max_seqlen=num_contextual_features, + backbone=network_args.backbone, + transformer_ffn_dim=network_args.transformer_ffn_dim, ) sm_major = torch.cuda.get_device_capability()[0] diff --git a/examples/hstu/inference/transformer.md b/examples/hstu/inference/transformer.md new file mode 100644 index 000000000..c5b30ec92 --- /dev/null +++ b/examples/hstu/inference/transformer.md @@ -0,0 +1,133 @@ +# Transformer recommendation inference + +Set `NetworkArgs.backbone = 'transformer'` to select a Transformer dense +backbone in the existing ranking inference workflow. The default remains +`'hstu'`. Sparse embeddings, item/action interleaving, context tokens, position +embeddings, candidate postprocessing, prediction heads and KV-cache management +use the existing recommendation components. + +## Architecture and input contract + +Each Transformer layer applies Pre-LayerNorm, biased Q/K/V projections, +scaled softmax multi-head self-attention, a biased output projection, and a +Pre-LayerNorm GELU feed-forward network. Both sublayers have residual additions +when `residual=True`. `transformer_ffn_dim` defaults to four times `hidden_size`. +The attention projection width is `num_heads * head_dim`; it need not equal +`hidden_size`. All heads have K/V, with matching K/V dimensions. There is no +dropout in inference. + +Position information comes from the existing preprocessing position embeddings. +This is not a Llama/Hugging Face checkpoint loader: RoPE, GQA/MQA, cross-attention +and model-specific weight conversion are not provided. Context tokens use the +causal inference mask, like history tokens; they do not attend to future history. +The existing training entry points remain HSTU-only and reject a Transformer +selection rather than silently training the wrong architecture. + +Packed inputs must contain each user's context/history followed by candidates. +History queries attend causally. A candidate sees the history and itself, but +not other candidates (`target_group_size=1`). A plain causal mask would introduce +cross-candidate dependencies and is not equivalent. `max_seq_len` bounds the +complete sequence, including candidates, even when the current query only +contains the uncached suffix. + +## Running + +In the repository's CUDA development environment, from `examples/hstu`: + +```bash +export PYTHONPATH="$(realpath ..):$PWD:$PYTHONPATH" +python inference/inference_gr_ranking.py \ + --gin_config_file inference/configs/kuairand_1k_transformer_ranking.gin \ + --checkpoint_dir /path/to/transformer-checkpoint --mode eval +``` + +The `get_inference_hstu_config` API also accepts `backbone="transformer"` and +`transformer_ffn_dim=...`; its historical name is retained for compatibility. +The Triton Python dense model reads the same NetworkArgs configuration. + +## Checkpoints + +The sparse checkpoint layout is unchanged. The dense state lives in the existing +`torch_module/model.0.pth` file under `model_state_dict`. Transformer layers use +these names beneath `_hstu_block._attention_layers..`: + +- `input_norm.weight`, `input_norm.bias` +- `qkv.weight`, `qkv.bias` (Q then K then V, each of width `num_heads * head_dim`) +- `proj.weight`, `proj.bias` +- `ffn_norm.weight`, `ffn_norm.bias` +- `ffn.0.weight`, `ffn.0.bias`, `ffn.2.weight`, `ffn.2.bias` + +The shared processor and MLP retain their existing names. To obtain the complete +dense schema, construct `get_inference_ranking_gr(...)` with a Transformer config +and inspect `model.dense_module.state_dict()`. Nonpersistent capture buffers are +excluded. An external trainer/converter must provide weights for this exact +architecture and the shared processor/head, plus the matching sparse weights. +Missing/unexpected keys fail loading, including when the outer workflow uses +`strict=False` to filter embedding keys. HSTU UVQK transposition and cached weight +refresh are applied only to HSTU checkpoints. + +## Cache, graph and export paths + +The layer uses the existing NHD page layout +`[pages, 2, page_size, num_heads, head_dim]`. New history K/V are appended through +`paged_kvcache_ops.append_kvcache` on CUDA; candidates are never persisted. +Reads materialize padded K/V and use PyTorch SDPA with an explicit recommendation +mask and cached query offset. A tensor implementation of the same cache writes +is available on CPU for numerical tests. It does not emulate asynchronous GPU +transfers or the cache manager. + +The layer implements `forward_naive`, `forward_input`, `forward_output` and +`output_buffer_` for the existing per-layer CUDA graph capture/replay orchestration. +Capture passes a nonzero token-count upper bound to the append operator so it +does not read a GPU scalar back to the host. Native cache onload synchronization +and offload are owned by the existing inference orchestration. + +Both existing exporter entry points propagate the backbone selection: + +```bash +python inference_aoti/export_inference_gr_ranking.py \ + --gin_config_file inference/configs/kuairand_1k_transformer_ranking.gin \ + --checkpoint_dir /path/to/transformer-checkpoint --max_bs 2 \ + --export_dir /path/to/empty-transformer-export \ + --dump_dir /path/to/empty-transformer-replay + +python inference_aoti/export_inference_gr_ranking_kvcache.py \ + --gin_config_file inference/configs/kuairand_1k_transformer_ranking.gin \ + --checkpoint_dir /path/to/transformer-checkpoint --max_bs 2 \ + --kvcache_config_file inference_aoti/kvcache_cpp_runtime.yaml \ + --export_dir /path/to/empty-transformer-kv-export \ + --dump_dir /path/to/empty-transformer-kv-replay +``` + +Adjust the KV runtime YAML dimensions, dtype and maximum lengths to match the +model. Follow the [AOTI workflow](../inference_aoti/README.md) for dependency +versions, C++ replay and Triton deployment. The exporters use the existing +training shell to obtain sparse/processor/head schemas; they construct fresh +Transformer layers and load the supplied Transformer checkpoint. They do not +convert HSTU dense weights into Transformer weights. + +## Validation and limitations + +```bash +python -m pytest test/test_transformer_inference.py \ + test/test_hstu_block_inference.py test/test_nve_aoti_compat.py -q -ra +``` + +CPU tests compare the actual layer against independent, unpadded softmax +arithmetic; cover multi-layer paged-prefix reuse, variable-length users, +noncontiguous physical pages, candidate independence, page-tail preservation, +empty histories, half/bfloat16, HSTU/Transformer checkpoint handling, and +torch.export save/reload with dynamic shapes and cache mutation. CUDA-only tests +exercise the native append operator, graph capture/replay and the shared +recommendation pre/postprocessor. A skip is not GPU validation. + +This backend prioritizes functional integration. Eager attention bounds query +padding by the smaller of the packed token count and `max_seq_len`; export uses +the fixed `max_seq_len` bound to avoid specializing dynamic token dimensions. +Cached K/V always pad to `max_seq_len`. Attention creates a dense boolean mask +per user, so large configured limits can consume substantial memory. This is not an +optimized paged Transformer attention kernel and makes no latency/throughput +claim. Profile realistic workloads before production use. This implementation +was locally tested on CPU; full NVE + GPU cache manager + AOTI/C++ + Triton +integration and recommendation-quality benchmarks still require a compatible +CUDA environment and a trained Transformer checkpoint. diff --git a/examples/hstu/inference/triton/hstu_model/model.py b/examples/hstu/inference/triton/hstu_model/model.py index a64c7af5b..2d6a0d59b 100644 --- a/examples/hstu/inference/triton/hstu_model/model.py +++ b/examples/hstu/inference/triton/hstu_model/model.py @@ -209,6 +209,8 @@ def get_inference_dense_model_with_feature_names( dtype=inference_dtype, position_encoding_config=position_encoding_config, contextual_max_seqlen=num_contextual_features, + backbone=network_args.backbone, + transformer_ffn_dim=network_args.transformer_ffn_dim, ) ranking_args = RankingArgs() diff --git a/examples/hstu/inference_aoti/export_inference_gr_ranking.py b/examples/hstu/inference_aoti/export_inference_gr_ranking.py index 339cb13be..b073cbff5 100644 --- a/examples/hstu/inference_aoti/export_inference_gr_ranking.py +++ b/examples/hstu/inference_aoti/export_inference_gr_ranking.py @@ -178,8 +178,66 @@ def get_exportable_model_for_inference( dynamic_table_configs, trained_emb_table_sizes, checkpoint_dir, + max_batch_size=1, + total_max_seqlen=8192, + num_contextual_features=0, ): + """Build and load the selected inference backbone for no-cache export. + + Reuse the training model's sparse and prediction-head schemas. Transformer + export creates fresh dense layers and requires a matching checkpoint; + batch and sequence limits bound its padded attention tensors. + """ model = get_training_gr_model() + if NetworkArgs().backbone == "transformer": + from configs import get_inference_hstu_config + from dynamicemb.exportable_tables import apply_inference_embedding_collection + from model.inference_ranking_gr import InferenceRankingGR + from modules.exportable_embedding import apply_inference_sparse + from modules.inference_dense_module import InferenceDenseModule + + # The training shell supplies the existing sparse/MLP/processor schemas. + # Its HSTU weights are never used; load a matching Transformer checkpoint. + model = apply_inference_embedding_collection( + model, dynamic_table_configs, trained_emb_table_sizes + ) + cfg = model._hstu_config + inference_config = get_inference_hstu_config( + hidden_size=cfg.hidden_size, + num_layers=cfg.num_layers, + num_attention_heads=cfg.num_attention_heads, + head_dim=cfg.kv_channels, + max_batch_size=max_batch_size, + max_seq_len=total_max_seqlen, + norm_epsilon=cfg.layernorm_epsilon, + dtype=torch.bfloat16 + if cfg.bf16 + else torch.float16 + if cfg.fp16 + else torch.float32, + learnable_input_layernorm=cfg.learnable_input_layernorm, + residual=cfg.residual, + is_causal=cfg.is_causal, + target_group_size=cfg.target_group_size, + position_encoding_config=cfg.position_encoding_config, + hstu_preprocessing_config=cfg.hstu_preprocessing_config, + contextual_max_seqlen=num_contextual_features, + backbone="transformer", + transformer_ffn_dim=NetworkArgs().transformer_ffn_dim, + export_mode=True, + ) + inference_model = InferenceRankingGR( + apply_inference_sparse(model._embedding_collection), + InferenceDenseModule( + inference_config, None, model._task_config, mlp=model._mlp + ), + ) + if cfg.bf16: + inference_model.bfloat16() + elif cfg.fp16: + inference_model.half() + inference_model.load_checkpoint(checkpoint_dir) + return inference_model.eval() inference_model = apply_inference( model, dynamic_table_configs=dynamic_table_configs, @@ -257,6 +315,9 @@ def strip_padding_batch(batch, unpadded_batch_size): dynamic_table_configs, trained_emb_table_sizes, checkpoint_dir, + max_batch_size, + total_max_seqlen, + num_contextual_features, ) eval_module = get_multi_event_metric_module( diff --git a/examples/hstu/inference_aoti/export_inference_gr_ranking_kvcache.py b/examples/hstu/inference_aoti/export_inference_gr_ranking_kvcache.py index c38f15119..be21c9ba1 100644 --- a/examples/hstu/inference_aoti/export_inference_gr_ranking_kvcache.py +++ b/examples/hstu/inference_aoti/export_inference_gr_ranking_kvcache.py @@ -277,6 +277,9 @@ def make_inference_hstu_config( contextual_max_seqlen=contextual_max_seqlen, scaling_seqlen=hstu_config.scaling_seqlen, export_mode=True, + backbone=NetworkArgs().backbone, + transformer_ffn_dim=NetworkArgs().transformer_ffn_dim, + hstu_preprocessing_config=hstu_config.hstu_preprocessing_config, ) diff --git a/examples/hstu/modules/hstu_block_inference.py b/examples/hstu/modules/hstu_block_inference.py index b7c5257b9..c0cfaa5b0 100644 --- a/examples/hstu/modules/hstu_block_inference.py +++ b/examples/hstu/modules/hstu_block_inference.py @@ -9,6 +9,7 @@ from modules.hstu_processor import HSTUBlockPostprocessor, HSTUBlockPreprocessor from modules.jagged_data import JaggedData from modules.paged_hstu_infer_layer import PagedHSTUInferLayer +from modules.transformer_infer_layer import TransformerInferLayer from torchrec.sparse.jagged_tensor import JaggedTensor @@ -177,9 +178,14 @@ def __init__( self._preprocessor = HSTUBlockPreprocessor(config, is_inference=True) self._postprocessor = HSTUBlockPostprocessor(is_inference=True) + layer_type = ( + TransformerInferLayer + if config.backbone == "transformer" + else PagedHSTUInferLayer + ) self._attention_layers = torch.nn.ModuleList( [ - PagedHSTUInferLayer(config, layer_idx) + layer_type(config, layer_idx) for layer_idx in range(self.config.num_layers) ] ) @@ -203,8 +209,14 @@ def forward( """ with torch.inference_mode(): jd = self._preprocessor(embeddings, batch) - for hstu_layer in self._attention_layers: - jd = hstu_layer(jd) + jd.values = self.predict( + batch.batch_size, + jd.values.shape[0], + jd.values, + jd, + None, + use_cudagraph=False, + ) return self._postprocessor(jd) def predict( diff --git a/examples/hstu/modules/inference_checkpoint.py b/examples/hstu/modules/inference_checkpoint.py new file mode 100644 index 000000000..3baf8a73b --- /dev/null +++ b/examples/hstu/modules/inference_checkpoint.py @@ -0,0 +1,49 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 +"""Dense checkpoint loading shared by inference backbones.""" + +import torch + + +def load_dense_state_dict(module, state_dict, *args, **kwargs): + """Load dense weights using the selected backbone's checkpoint layout. + + Filter sparse embedding keys and apply legacy HSTU transpositions only to + HSTU layers. Forward load options to PyTorch and return its incompatible-key + result. Missing or unexpected dense keys raise RuntimeError even when the + caller passes strict=False to allow filtering the sparse weights. + """ + hstu_layout = not module._use_exportable and module._backbone == "hstu" + converted = {} + for key, value in state_dict.items(): + if ( + key.startswith( + "_embedding_collection._data_parallel_embedding_collection.embeddings." + ) + or "_model_parallel_embedding_collection" in key + ): + continue + new_key = key + if hstu_layout: + for old, new, transpose in ( + ("_linear_uvqk_weight", "_linear_uvqk.weight", True), + ("_linear_uvqk_bias", "_linear_uvqk.bias", False), + ("_linear_proj_weight", "_linear_proj.weight", True), + ): + if key.endswith(old): + new_key = key.removesuffix(old) + new + value = value.T if transpose else value + break + converted[new_key] = value + result = torch.nn.Module.load_state_dict(module, converted, *args, **kwargs) + if result.missing_keys or result.unexpected_keys: + raise RuntimeError( + f"Checkpoint does not match {module._backbone} backbone: " + f"missing={result.missing_keys}, unexpected={result.unexpected_keys}" + ) + if hstu_layout: + with torch.no_grad(): + for layer in module._hstu_block._attention_layers: + layer._linear_uvqk_weight.copy_(layer._linear_uvqk.weight.T) + layer._linear_proj_weight.copy_(layer._linear_proj.weight.T) + return result diff --git a/examples/hstu/modules/inference_dense_module.py b/examples/hstu/modules/inference_dense_module.py index b33951ea6..000253569 100755 --- a/examples/hstu/modules/inference_dense_module.py +++ b/examples/hstu/modules/inference_dense_module.py @@ -140,6 +140,7 @@ def __init__( self._hstu_config = hstu_config self._task_config = task_config self._use_exportable = use_exportable + self._backbone = getattr(hstu_config, "backbone", "hstu") if self._use_exportable: assert isinstance( hstu_config, HSTUConfig @@ -293,41 +294,10 @@ def load_checkpoint(self, checkpoint_dir): self.load_state_dict(model_state_dict, strict=False) def load_state_dict(self, model_state_dict, *args, **kwargs): - new_state_dict = {} - for k in model_state_dict: - if ( - k.startswith( - "_embedding_collection._data_parallel_embedding_collection.embeddings." - ) - or "_model_parallel_embedding_collection" in k - ): - continue - - is_transposed = False - - newk = k - if not self._use_exportable: - if k.endswith("_linear_uvqk_weight"): - newk = k.removesuffix("_linear_uvqk_weight") + "_linear_uvqk.weight" - is_transposed = True - elif k.endswith("_linear_uvqk_bias"): - newk = k.removesuffix("_linear_uvqk_bias") + "_linear_uvqk.bias" - elif k.endswith("_linear_proj_weight"): - newk = k.removesuffix("_linear_proj_weight") + "_linear_proj.weight" - is_transposed = True - - new_state_dict[newk] = ( - model_state_dict[k] if not is_transposed else model_state_dict[k].T - ) + """Load dense weights and refresh backbone-specific inference state.""" + from modules.inference_checkpoint import load_dense_state_dict - unloaded_modules = super().load_state_dict(new_state_dict, *args, **kwargs) - if not self._use_exportable: - for hstu_layer in self._hstu_block._attention_layers: - hstu_layer._linear_uvqk_weight.copy_(hstu_layer._linear_uvqk.weight.T) - hstu_layer._linear_proj_weight.copy_(hstu_layer._linear_proj.weight.T) - - assert unloaded_modules.missing_keys == [] - assert unloaded_modules.unexpected_keys == [] + return load_dense_state_dict(self, model_state_dict, *args, **kwargs) def forward_with_kvcache( self, @@ -455,6 +425,8 @@ def forward( embeddings: Dict[str, JaggedTensor], ): with torch.inference_mode(): + if not self._use_exportable: + return self.forward_nokvcache(batch, embeddings) # Forward through HSTU block jd_output, _ = self._hstu_block(embeddings, batch) # Prediction head diff --git a/examples/hstu/modules/transformer_infer_layer.py b/examples/hstu/modules/transformer_infer_layer.py new file mode 100644 index 000000000..29e288679 --- /dev/null +++ b/examples/hstu/modules/transformer_infer_layer.py @@ -0,0 +1,329 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 +"""Recommendation Transformer with the paged HSTU inference layer contract. + +The attention implementation materializes bounded, padded K/V tensors and uses +PyTorch SDPA. It is a correctness-oriented backend, not a fused paged-attention +kernel. Position embeddings are supplied by the shared recommendation processor. +This module deliberately imports only PyTorch so its numerical path also runs +on CPU. CUDA cache writes use the existing paged_kvcache_ops operator. +""" + +import torch +from torch import nn +from torch.nn import functional as F + + +def _pad(values, offsets, width): + lengths = offsets[1:] - offsets[:-1] + positions = torch.arange(width, device=values.device) + valid = positions[None, :] < lengths[:, None] + indices = torch.where(valid, offsets[:-1, None] + positions, values.shape[0]) + # A zero sentinel also handles an empty packed tensor without a negative index. + extended = torch.cat((values, values.new_zeros((1,) + values.shape[1:]))) + return extended[indices.long()], valid + + +def _unpad(values, offsets, token_count): + tokens = torch.arange(token_count, device=values.device) + users = torch.searchsorted(offsets[1:].contiguous(), tokens, right=True) + users = users.clamp(max=offsets.shape[0] - 2) + positions = (tokens - offsets[users]).clamp(max=values.shape[1] - 1) + result = values[users, positions] + # CUDA graph buckets may contain unused packed-token slots. + return torch.where((tokens < offsets[-1])[:, None], result, 0) + + +def read_paged_history(table, page_ids, page_indptr, history_lengths, width): + """Gather NHD pages, masking every inactive slot before using its value.""" + positions = torch.arange(width, device=table.device) + valid = positions[None, :] < history_lengths[:, None] + logical_pages = page_indptr[:-1, None] + positions // table.shape[2] + # Zero-history users need no pages; never dereference their inactive metadata. + ids = torch.cat((page_ids, page_ids.new_zeros(1))) + logical_pages = torch.where(valid, logical_pages, page_ids.shape[0]) + physical_pages = ids[logical_pages.long()].long() + page_offsets = positions % table.shape[2] + key = table[physical_pages, 0, page_offsets] + value = table[physical_pages, 1, page_offsets] + live = valid[:, :, None, None] + return torch.where(live, key, 0), torch.where(live, value, 0) + + +class TransformerInferLayer(nn.Module): + """Apply Pre-LN MHA and a GELU FFN with independent candidates. + + Args: + config: Inference configuration defining dimensions, dtype and limits. + layer_idx: Index into the shared per-layer KV-cache table collection. + device: Parameter device; defaults to the current CUDA device. + + Attributes: + num_heads: Public attention head count for tensor layout inspection. + head_dim: Public per-head dimension for tensor layout inspection. + input_norm: Input normalization module in the checkpoint schema. + qkv: Fused Q/K/V projection module in the checkpoint schema. + proj: Attention output projection module in the checkpoint schema. + ffn_norm: Feed-forward normalization module in the checkpoint schema. + ffn: Feed-forward modules in the checkpoint schema. + output_buffer_: Output storage required by the shared graph runner. + """ + + def __init__(self, config, layer_idx, device=None): + super().__init__() + self._layer_idx = layer_idx + self.num_heads = config.num_heads + self.head_dim = config.head_dim + self._max_seq_len = config.max_seq_len + self._export_mode = config.export_mode + self._residual = config.residual + dtype = ( + torch.bfloat16 + if config.bf16 + else torch.float16 + if config.fp16 + else torch.float32 + ) + if device is None: + device = torch.device("cuda", torch.cuda.current_device()) + kwargs = dict(device=device, dtype=dtype) + inner = self.num_heads * self.head_dim + self.input_norm = nn.LayerNorm( + config.hidden_size, + eps=config.layernorm_epsilon, + elementwise_affine=getattr(config, "learnable_input_layernorm", True), + **kwargs, + ) + self.qkv = nn.Linear(config.hidden_size, 3 * inner, **kwargs) + self.proj = nn.Linear(inner, config.hidden_size, **kwargs) + self.ffn_norm = nn.LayerNorm( + config.hidden_size, eps=config.layernorm_epsilon, **kwargs + ) + ffn_size = config.transformer_ffn_dim or 4 * config.hidden_size + self.ffn = nn.Sequential( + nn.Linear(config.hidden_size, ffn_size, **kwargs), + nn.GELU(), + nn.Linear(ffn_size, config.hidden_size, **kwargs), + ) + capacity = config.max_batch_size * config.max_seq_len + self.register_buffer( + "_qkv_buffer", torch.empty(capacity, 3 * inner, **kwargs), persistent=False + ) + self.register_buffer( + "output_buffer_", + torch.empty(capacity, config.hidden_size, **kwargs), + persistent=False, + ) + self.requires_grad_(False) + + def project_qkv(self, hidden): + """Project packed hidden states to Q/K/V shaped as tokens, heads, dim.""" + mixed = self.qkv(self.input_norm(hidden)) + return tuple( + x.reshape(-1, self.num_heads, self.head_dim) for x in mixed.chunk(3, dim=-1) + ) + + def attention( + self, + query, + key, + value, + offsets, + candidates, + cache_table=None, + page_ids=None, + page_indptr=None, + history_lengths=None, + ): + """Attend to causal history and each candidate's own position. + + Offsets delimit packed users; candidates give their trailing candidate + counts. When a cache table is supplied, history_lengths includes the + cached prefix and newly appended history. Return packed attention + outputs with the head dimensions flattened. + """ + if query.shape[0] == 0: + return query.new_empty((0, self.num_heads * self.head_dim)) + # A cached suffix often contains far fewer tokens than the configured + # maximum. Shape bounds avoid a GPU-to-host maximum-length reduction. + # A fixed export bound permits token dimensions to cross max_seq_len + # (the packed batch can exceed one user's limit) without shape guards. + query_width = ( + self._max_seq_len + if torch.compiler.is_compiling() + else min(self._max_seq_len, query.shape[0]) + ) + q, query_valid = _pad(query, offsets, query_width) + k, _ = _pad(key, offsets, query_width) + v, _ = _pad(value, offsets, query_width) + lengths = offsets[1:] - offsets[:-1] + torch._assert_async( + torch.all((lengths >= 0) & (lengths <= self._max_seq_len)), + "sequence length exceeds Transformer max_seq_len", + ) + torch._assert_async( + torch.all((candidates >= 0) & (candidates <= lengths)), + "invalid Transformer candidate count", + ) + new_history = lengths - candidates + query_local_positions = torch.arange(query_width, device=query.device) + key_width = self._max_seq_len if cache_table is not None else query_width + positions = torch.arange(key_width, device=query.device) + if cache_table is None: + history_lengths = new_history + cached_lengths = torch.zeros_like(lengths) + else: + cached_lengths = history_lengths - new_history + torch._assert_async( + torch.all(cached_lengths >= 0), "negative cached prefix length" + ) + torch._assert_async( + torch.all(history_lengths + candidates <= self._max_seq_len), + "cached sequence exceeds Transformer max_seq_len", + ) + history_k, history_v = read_paged_history( + cache_table, page_ids, page_indptr, history_lengths, self._max_seq_len + ) + # Candidates are not persisted in the cache. Read them from this call. + local_positions = ( + (positions[None, :] - cached_lengths[:, None]) + .clamp(min=0, max=query_width - 1) + .long() + ) + users = torch.arange(lengths.shape[0], device=query.device)[:, None] + use_history = (positions[None, :] < history_lengths[:, None])[ + :, :, None, None + ] + k = torch.where(use_history, history_k, k[users, local_positions]) + v = torch.where(use_history, history_v, v[users, local_positions]) + query_positions = query_local_positions[None, :] + cached_lengths[:, None] + keys = positions[None, None, :] + queries = query_positions[:, :, None] + # Histories are causal, candidates see history and themselves only. + allowed = ( + (keys <= queries) + & ((keys < history_lengths[:, None, None]) | (keys == queries)) + & (keys < (history_lengths + candidates)[:, None, None]) + & query_valid[:, :, None] + ) + output = F.scaled_dot_product_attention( + q.transpose(1, 2), + k.transpose(1, 2), + v.transpose(1, 2), + attn_mask=allowed[:, None], + dropout_p=0.0, + ) + output = output.transpose(1, 2).flatten(2) + return _unpad(output, offsets, query.shape[0]) + + def finish(self, hidden, attention): + """Apply the output projection, feed-forward block and residuals.""" + projected = self.proj(attention) + hidden = hidden + projected if self._residual else projected + output = self.ffn(self.ffn_norm(hidden)) + return hidden + output if self._residual else output + + def forward(self, hidden, offsets, candidates): + """Tensor-only no-cache entry point, also usable with torch.export.""" + q, k, v = self.project_qkv(hidden) + return self.finish(hidden, self.attention(q, k, v, offsets, candidates)) + + def _append(self, key, value, offsets, candidates, metadata, batch_size): + table = metadata.kv_cache_table[self._layer_idx] + if table.is_cuda: + return torch.ops.paged_kvcache_ops.append_kvcache( + key, + value, + metadata.batch_indices, + metadata.position, + torch.cat((candidates.new_zeros(1), candidates.cumsum(0))).to( + offsets.dtype + ), + # A positive upper bound avoids the operator's nnz.item() path + # during CUDA graph capture. The kernel reads the live count. + metadata.new_history_nnz_cuda, + 0 if self._export_mode else key.shape[0], + table, + metadata.kv_indices, + metadata.kv_indptr, + metadata.kv_last_page_len, + 0, + ) + # CPU reference backend: same persistent NHD format, history tokens only. + tokens = torch.arange(key.shape[0], device=key.device) + users = torch.searchsorted(offsets[1:].contiguous(), tokens, right=True).clamp( + max=batch_size - 1 + ) + local = tokens - offsets[users] + lengths = offsets[1:] - offsets[:-1] + new_history = lengths - candidates + valid = (local < new_history[users]) & (tokens < offsets[-1]) + users, local = users[valid], local[valid] + positions = metadata.total_history_lengths[users] - new_history[users] + local + pages = metadata.kv_indices[ + (metadata.kv_indptr[users] + positions // table.shape[2]).long() + ].long() + table[pages, 0, positions % table.shape[2]] = key[valid] + table[pages, 1, positions % table.shape[2]] = value[valid] + return table + + def _compute(self, hidden, q, k, v, jd, metadata, batch_size, append): + offsets = jd.seqlen_offsets[: batch_size + 1] + candidates = ( + jd.num_candidates[:batch_size] + if jd.num_candidates is not None + else torch.zeros_like(offsets[:-1]) + ) + cache = {} + if metadata is not None: + table = ( + self._append(k, v, offsets, candidates, metadata, batch_size) + if append + else metadata.kv_cache_table[self._layer_idx] + ) + handle = metadata.kv_onload_handle + if not self._export_mode and handle is not None: + handle.stream_wait_layer(self._layer_idx) + cache = dict( + cache_table=table, + page_ids=metadata.kv_indices, + page_indptr=metadata.kv_indptr[: batch_size + 1], + history_lengths=metadata.total_history_lengths[:batch_size], + ) + attended = self.attention(q, k, v, offsets, candidates, **cache) + return self.finish(hidden, attended) + + def forward_naive(self, batch_size, num_tokens, hidden, jd, metadata): + """Run one layer, appending history when cache metadata is supplied.""" + hidden = hidden[:num_tokens] + q, k, v = self.project_qkv(hidden) + return self._compute(hidden, q, k, v, jd, metadata, batch_size, append=True) + + def forward_input(self, batch_size, num_tokens, hidden, jd, metadata): + """Store projected Q/K/V and append history for split graph execution.""" + mixed = self.qkv(self.input_norm(hidden[:num_tokens])) + self._qkv_buffer[:num_tokens].copy_(mixed) + if metadata is not None: + _, k, v = ( + x.reshape(-1, self.num_heads, self.head_dim) for x in mixed.chunk(3, -1) + ) + offsets = jd.seqlen_offsets[: batch_size + 1] + candidates = ( + jd.num_candidates[:batch_size] + if jd.num_candidates is not None + else torch.zeros_like(offsets[:-1]) + ) + self._append(k, v, offsets, candidates, metadata, batch_size) + return self._qkv_buffer[:num_tokens] + + def forward_output(self, batch_size, num_tokens, hidden, jd, metadata): + """Attend with stored Q/K/V and return the shared output buffer view.""" + q, k, v = ( + x.reshape(-1, self.num_heads, self.head_dim) + for x in self._qkv_buffer[:num_tokens].chunk(3, -1) + ) + output = self._compute( + hidden[:num_tokens], q, k, v, jd, metadata, batch_size, append=False + ) + self.output_buffer_[:num_tokens].copy_(output) + return self.output_buffer_[:num_tokens] diff --git a/examples/hstu/test/test_hstu_block_inference.py b/examples/hstu/test/test_hstu_block_inference.py index 11fbb5eda..394be3529 100755 --- a/examples/hstu/test/test_hstu_block_inference.py +++ b/examples/hstu/test/test_hstu_block_inference.py @@ -13,221 +13,119 @@ # See the License for the specific language governing permissions and # limitations under the License. import sys +from pathlib import Path +import pytest import torch -from commons.datasets.hstu_batch import FeatureConfig -from commons.datasets.random_inference_dataset import RandomInferenceDataGenerator -from configs import ( - InferenceEmbeddingConfig, - RankingConfig, - get_inference_hstu_config, - get_kvcache_config, -) - -sys.path.append("./model/") -from inference_ranking_gr import InferenceRankingGR - - -def get_test_setup(): - max_batch_size = 2 - max_seqlen = 1024 +import torch.nn.functional as F - # context_emb_size = 1000 - item_fea_name, item_vocab_size = "item_feat", 10000 - action_fea_name, action_vocab_size = "act_feat", 128 - feature_configs = [ - FeatureConfig( - feature_names=[item_fea_name, action_fea_name], - max_item_ids=[item_vocab_size - 1, action_vocab_size - 1], - max_sequence_length=max_seqlen, - is_jagged=False, - ), - ] - max_contextual_seqlen = 0 +HSTU_ROOT = Path(__file__).resolve().parents[1] +for path in (HSTU_ROOT, HSTU_ROOT.parent): + sys.path.insert(0, str(path)) - hidden_dim_size = 512 - num_heads = 4 - head_dim = 128 - num_layers = 4 - hstu_config = get_inference_hstu_config( - hidden_dim_size, - num_layers, - num_heads, - head_dim, - max_batch_size, - max_seqlen, +@pytest.mark.skipif( + not torch.cuda.is_available(), + reason="requires the CUDA recommendation preprocessing operators", +) +@pytest.mark.parametrize("backbone", ["hstu", "transformer"]) +@torch.inference_mode() +def test_hstu_process_inference(backbone): + # Exercise the current processor/block interfaces directly. The old test + # used a removed RandomInferenceDataGenerator and obsolete cache APIs even + # though it was testing only preprocessing and candidate extraction. + from commons.datasets.hstu_batch import HSTUBatch + from configs import get_inference_hstu_config + from modules.hstu_block_inference import HSTUBlockInference + from modules.transformer_infer_layer import TransformerInferLayer + from torchrec.sparse.jagged_tensor import JaggedTensor, KeyedJaggedTensor + + device = torch.device("cuda") + cfg = get_inference_hstu_config( + hidden_size=128, + num_layers=1, + num_attention_heads=2, + head_dim=64, + max_batch_size=2, + max_seq_len=32, + dtype=torch.float32, + contextual_max_seqlen=1, + backbone=backbone, ) - - _blocks_in_primary_pool = 10240 - _page_size = 32 - _offload_chunksize = 128 - kv_cache_config = get_kvcache_config( - blocks_in_primary_pool=_blocks_in_primary_pool, - page_size=_page_size, - offload_chunksize=_offload_chunksize, + block = HSTUBlockInference(cfg).to(device) + assert isinstance(block._attention_layers[0], TransformerInferLayer) == ( + backbone == "transformer" ) - emb_configs = [ - InferenceEmbeddingConfig( - feature_names=["act_feat"], - table_name="act", - vocab_size=action_vocab_size, - dim=hidden_dim_size, - use_dynamicemb=False, - ), - InferenceEmbeddingConfig( - feature_names=["context_feat", "item_feat"] - if max_contextual_seqlen > 0 - else ["item_feat"], - table_name="item", - vocab_size=item_vocab_size, - dim=hidden_dim_size, - use_dynamicemb=True, - ), - ] - num_tasks = 3 - task_config = RankingConfig( - embedding_configs=emb_configs, - prediction_head_arch=[[128, 10, 1] for _ in range(num_tasks)], + row_lengths = {"context": [1, 1], "item": [3, 3], "action": [2, 1]} + embeddings = {} + for name, lengths in row_lengths.items(): + embeddings[name] = JaggedTensor( + values=torch.randn(sum(lengths), 128, device=device), + lengths=torch.tensor(lengths, device=device), + ) + features = KeyedJaggedTensor( + keys=list(row_lengths), + values=torch.arange(11, device=device), + lengths=torch.tensor([1, 1, 3, 3, 2, 1], device=device), ) - - model = InferenceRankingGR( - hstu_config=hstu_config, - kvcache_config=kv_cache_config, - task_config=task_config, - use_cudagraph=False, + batch = HSTUBatch( + features=features, + batch_size=2, + # The shared schema bounds include candidate slots for item and action, + # although inference action values themselves contain history only. + feature_to_max_seqlen={"context": 1, "item": 4, "action": 4}, + contextual_feature_names=["context"], + item_feature_name="item", + action_feature_name="action", + max_num_candidates=2, + num_candidates=torch.tensor([1, 2], device=device), ) - model.bfloat16() - model.eval() - - return model, feature_configs - - -def test_hstu_process_inference(): - max_batch_size = 2 - max_seqlen = 1024 - max_num_candidates = 128 - max_incremental_seqlen = 64 - - item_fea_name = "item_feat" - action_fea_name = "act_feat" - - with torch.inference_mode(): - model_predict, feature_configs = get_test_setup() - - data_generator = RandomInferenceDataGenerator( - feature_configs, - item_fea_name, - [], - action_fea_name, - 1024, - max_batch_size, - max_seqlen, - max_num_candidates, - max_incremental_seqlen, - False, + jd = block._preprocessor(embeddings, batch) + ctx, items, actions = ( + embeddings[k].values() for k in ("context", "item", "action") + ) + expected = torch.stack( + [ + ctx[0], + items[0], + actions[0], + items[1], + actions[1], + items[2], + ctx[1], + items[3], + actions[2], + items[4], + items[5], + ] + ) + torch.testing.assert_close(jd.values, expected) + torch.testing.assert_close( + jd.seqlen, torch.tensor([6, 5], dtype=torch.int32, device=device) + ) + post = block._postprocessor(jd) + expected_candidates = F.normalize( + items[torch.tensor([2, 4, 5], device=device)], dim=-1, eps=1e-6 + ) + torch.testing.assert_close(post.values, expected_candidates) + if backbone == "transformer": + output = block(embeddings, batch) + assert output.values.shape == (3, 128) + assert torch.isfinite(output.values).all() + from configs import InferenceEmbeddingConfig, RankingConfig + from modules.inference_dense_module import InferenceDenseModule + + task = RankingConfig( + embedding_configs=[ + InferenceEmbeddingConfig(["item"], "item", 16, 128, False) + ], + prediction_head_arch=[128, 2], + num_tasks=2, ) - - num_test_batches = 100 - - for idx in range(num_test_batches): - uids = data_generator.get_inference_batch_user_ids() - - cached_start_pos, cached_len = model_predict.get_user_kvdata_info(uids) - truncate_start_pos = cached_start_pos + cached_len - - batch = data_generator.get_random_inference_batch(uids, truncate_start_pos) - - kvc_mtdt = model_predict.prepare_kv_cache(batch, uids, truncate_start_pos) - - embs = model_predict._embedding_collection(batch.features) - - jd = model_predict._hstu_block._preprocessor(embs, batch) - - history_lens = [ - (jd.seqlen[i].item() - jd.num_candidates[i].item()) // 2 - for i in range(batch.batch_size) - ] - - original_items = torch.tensor( - [ - embs["item_feat"].offsets()[i].item() + token_idx - for i in range(batch.batch_size) - for token_idx in range(history_lens[i]) - ] - ).long() - new_items = torch.tensor( - [ - 2 * token_idx + jd.seqlen_offsets[i].item() - for i in range(batch.batch_size) - for token_idx in range(history_lens[i]) - ] - ).long() - assert torch.allclose( - embs["item_feat"].values()[original_items].to(torch.bfloat16), - jd.values[new_items], - ) - - original_actions = torch.tensor( - [ - embs["act_feat"].offsets()[i].item() + token_idx - for i in range(batch.batch_size) - for token_idx in range(history_lens[i]) - ] - ).long() - new_actions = torch.tensor( - [ - 2 * token_idx + 1 + jd.seqlen_offsets[i].item() - for i in range(batch.batch_size) - for token_idx in range(history_lens[i]) - ] - ).long() - assert torch.allclose( - embs["act_feat"].values()[original_actions].to(torch.bfloat16), - jd.values[new_actions], - ) - - original_candidates = torch.tensor( - [ - embs["item_feat"].offsets()[i].item() + history_lens[i] + token_idx - for i in range(batch.batch_size) - for token_idx in range(jd.num_candidates[i].item()) - ] - ).long() - new_candidates = torch.tensor( - [ - jd.seqlen_offsets[i].item() + history_lens[i] * 2 + token_idx - for i in range(batch.batch_size) - for token_idx in range(jd.num_candidates[i].item()) - ] - ).long() - assert torch.allclose( - embs["item_feat"].values()[original_candidates].to(torch.bfloat16), - jd.values[new_candidates], - ) - - # post process - post_jd = model_predict._hstu_block._postprocessor(jd) - original_candidates = torch.tensor( - [ - embs["item_feat"].offsets()[i].item() + history_lens[i] + token_idx - for i in range(batch.batch_size) - for token_idx in range(jd.num_candidates[i].item()) - ] - ).long() - new_candidates = torch.tensor( - [ - post_jd.seqlen_offsets[i].item() + token_idx - for i in range(batch.batch_size) - for token_idx in range(jd.num_candidates[i].item()) - ] - ).long() - post_embs = ( - embs["item_feat"].values()[original_candidates].to(torch.bfloat16) - ) - post_embs = post_embs / torch.linalg.norm( - post_embs, ord=2, dim=-1, keepdim=True - ).clamp(min=1e-6) - assert torch.allclose(post_embs, post_jd.values) - - model_predict.offload_kv_cache(uids, kvc_mtdt) + dense = InferenceDenseModule(cfg, None, task, hstu_block=block).eval() + logits = dense(batch, embeddings) + expected_logits = dense._mlp(output.values) + torch.testing.assert_close(logits, expected_logits) + state = {k: v.clone() for k, v in dense.state_dict().items()} + dense.load_state_dict(state, strict=True) + torch.testing.assert_close(dense(batch, embeddings), logits) diff --git a/examples/hstu/test/test_transformer_inference.py b/examples/hstu/test/test_transformer_inference.py new file mode 100644 index 000000000..7dc9c8d68 --- /dev/null +++ b/examples/hstu/test/test_transformer_inference.py @@ -0,0 +1,506 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 +"""CPU numerical tests for the production Transformer inference implementation. + +No CUDA modules, monkeypatched operators or extracted AST are used. The manual +reference computes each user's unpadded attention independently with softmax. +GPU-only integration is covered separately at the end of this file. +""" + +import copy +import sys +from dataclasses import dataclass +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +from torch import nn +from torch.nn import functional as F + +HSTU_ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(HSTU_ROOT)) + +from modules.inference_checkpoint import load_dense_state_dict +from modules.transformer_infer_layer import TransformerInferLayer + + +def config(**overrides): + args = dict( + num_heads=2, + head_dim=4, + hidden_size=12, + max_seq_len=24, + max_batch_size=4, + export_mode=False, + residual=True, + bf16=False, + fp16=False, + layernorm_epsilon=1e-5, + transformer_ffn_dim=20, + ) + args.update(overrides) + return SimpleNamespace(**args) + + +def test_backbone_configuration_defaults_and_validation(): + from configs import get_inference_hstu_config + + args = dict( + hidden_size=12, + num_layers=2, + num_attention_heads=2, + head_dim=4, + max_batch_size=4, + max_seq_len=24, + ) + assert get_inference_hstu_config(**args).backbone == "hstu" + cfg = get_inference_hstu_config( + **args, backbone="transformer", transformer_ffn_dim=20, dtype=torch.float32 + ) + assert cfg.hstu_preprocessing_config is None + layer = TransformerInferLayer(cfg, 0, "cpu") + assert layer.ffn[0].out_features == 20 + with pytest.raises(ValueError, match="Unknown inference backbone"): + get_inference_hstu_config(**args, backbone="typo") + with pytest.raises(ValueError, match="positive"): + get_inference_hstu_config(**args, backbone="transformer", transformer_ffn_dim=0) + + +def offsets(lengths): + return torch.tensor( + [0] + list(torch.tensor(lengths).cumsum(0).tolist()), dtype=torch.int32 + ) + + +def reference(layer, x, candidates): + """Independent unpadded Pre-LN Transformer arithmetic for one user.""" + normed = F.layer_norm( + x, + (x.shape[-1],), + layer.input_norm.weight, + layer.input_norm.bias, + layer.input_norm.eps, + ) + q, k, v = F.linear(normed, layer.qkv.weight, layer.qkv.bias).chunk(3, -1) + q, k, v = [ + z.reshape(-1, layer.num_heads, layer.head_dim).transpose(0, 1) + for z in (q, k, v) + ] + n = x.shape[0] + allowed = torch.tensor( + [ + [j <= i and (j < n - candidates or j == i) for j in range(n)] + for i in range(n) + ], + device=x.device, + ) + logits = (q @ k.transpose(-1, -2)) / layer.head_dim**0.5 + weights = logits.masked_fill(~allowed, -float("inf")).softmax(-1) + out = (weights @ v).transpose(0, 1).reshape(n, -1) + out = F.linear(out, layer.proj.weight, layer.proj.bias) + x = x + out if layer._residual else out + normed = F.layer_norm( + x, + (x.shape[-1],), + layer.ffn_norm.weight, + layer.ffn_norm.bias, + layer.ffn_norm.eps, + ) + ffn = F.linear( + F.gelu(F.linear(normed, layer.ffn[0].weight, layer.ffn[0].bias)), + layer.ffn[2].weight, + layer.ffn[2].bias, + ) + return x + ffn if layer._residual else ffn + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) +@pytest.mark.parametrize("residual", [True, False]) +@pytest.mark.parametrize("seed", range(5)) +def test_ragged_forward_matches_unpadded_reference(dtype, residual, seed): + torch.manual_seed(seed) + layer = TransformerInferLayer(config(residual=residual), 0, "cpu").to(dtype) + lengths, candidates = [1, 5, 11, 20], [1, 2, 0, 4] + x = torch.randn(sum(lengths), 12, dtype=dtype) + expected = torch.cat( + [reference(layer, user, c) for user, c in zip(x.split(lengths), candidates)] + ) + actual = layer(x, offsets(lengths), torch.tensor(candidates)) + torch.testing.assert_close(actual, expected, atol=1e-6, rtol=2e-5) + + +def make_metadata(histories, layers, heads, dim, dtype, page_size=4): + counts = [(h + page_size - 1) // page_size for h in histories] + pages = sum(counts) + return SimpleNamespace( + kv_cache_table=[ + torch.full( + (max(1, pages), 2, page_size, heads, dim), float("nan"), dtype=dtype + ) + for _ in range(layers) + ], + kv_indices=torch.randperm(pages).to(torch.int32), + kv_indptr=offsets(counts), + total_history_lengths=torch.tensor(histories, dtype=torch.int32), + kv_onload_handle=None, + ) + + +def jagged(lengths, candidates): + return SimpleNamespace( + seqlen_offsets=offsets(lengths), num_candidates=torch.tensor(candidates) + ) + + +@pytest.mark.parametrize("seed", range(8)) +@pytest.mark.parametrize("page_size", [1, 4, 8]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) +@pytest.mark.parametrize("split_path", [False, True]) +@torch.no_grad() +def test_two_layer_paged_cache_matches_full_reference( + seed, page_size, dtype, split_path +): + torch.manual_seed(seed) + layers = nn.ModuleList( + [TransformerInferLayer(config(), i, "cpu").to(dtype) for i in range(2)] + ) + histories, candidates = [0, 3, 8, 16], [1, 2, 3, 4] + prefixes = [0, seed % 4, seed % 9, 16 if seed % 2 else 7] + lengths = [h + c for h, c in zip(histories, candidates)] + users = [torch.randn(n, 12, dtype=dtype) for n in lengths] + full = users + for layer in layers: + full = [reference(layer, x, c) for x, c in zip(full, candidates)] + metadata = make_metadata(histories, 2, 2, 4, dtype, page_size) + + # Populate cache using independent full-history projections, layer by layer. + # Prefix K/V are invariant to future history with the causal history mask. + history_inputs = [x[:h] for x, h in zip(users, histories)] + for i, layer in enumerate(layers): + next_history = [] + for u, (x, prefix) in enumerate(zip(history_inputs, prefixes)): + if x.shape[0]: + q, k, v = layer.project_qkv(x) + for p in range(prefix): + page = metadata.kv_indices[metadata.kv_indptr[u] + p // page_size] + metadata.kv_cache_table[i][page, 0, p % page_size] = k[p] + metadata.kv_cache_table[i][page, 1, p % page_size] = v[p] + next_history.append(reference(layer, x, 0)) + else: + next_history.append(x) + history_inputs = next_history + + delta_lengths = [n - p for n, p in zip(lengths, prefixes)] + jd = jagged(delta_lengths, candidates) + packed = torch.cat([x[p:] for x, p in zip(users, prefixes)]) + for layer in layers: + if split_path: + layer.forward_input(4, packed.shape[0], packed, jd, metadata) + packed = layer.forward_output( + 4, packed.shape[0], packed, jd, metadata + ).clone() + else: + packed = layer.forward_naive(4, packed.shape[0], packed, jd, metadata) + expected = torch.cat([x[p:] for x, p in zip(full, prefixes)]) + torch.testing.assert_close(packed, expected, atol=2e-6, rtol=2e-5) + assert torch.isfinite(packed).all() + # Candidate slots and unused page tails must never be persisted. + for table in metadata.kv_cache_table: + for u, h in enumerate(histories): + if h and h % page_size: + last_page = metadata.kv_indices[metadata.kv_indptr[u + 1] - 1] + assert torch.isnan(table[last_page, :, h % page_size :]).all() + + +def test_candidates_are_independent_and_permutation_equivariant(): + torch.manual_seed(9) + layer = TransformerInferLayer(config(), 0, "cpu").double() + x = torch.randn(9, 12, dtype=torch.float64) + off, c = offsets([9]), torch.tensor([3]) + before = layer(x, off, c) + changed = x.clone() + changed[6] += torch.arange(12) + after = layer(changed, off, c) + torch.testing.assert_close(before[7:], after[7:], rtol=0, atol=0) + perm = torch.tensor([0, 1, 2, 3, 4, 5, 8, 6, 7]) + torch.testing.assert_close(layer(x[perm], off, c), before[perm]) + + +def test_empty_history_and_empty_users(): + layer = TransformerInferLayer(config(), 0, "cpu") + x = torch.randn(3, 12) + jd = jagged([0, 2, 1], [0, 2, 1]) + metadata = make_metadata([0, 0, 0], 1, 2, 4, torch.float32) + got = layer.forward_naive(3, 3, x, jd, metadata) + torch.testing.assert_close(got, layer(x, jd.seqlen_offsets, jd.num_candidates)) + assert torch.isnan(metadata.kv_cache_table[0]).all() + + +def test_padded_packed_tail_is_ignored(): + layer = TransformerInferLayer(config(), 0, "cpu") + x = torch.randn(7, 12) + jd = jagged([3, 4], [1, 2]) + padded = torch.cat((x, torch.randn(9, 12))) + actual = layer.forward_naive(2, 16, padded, jd, None) + torch.testing.assert_close( + actual[:7], layer(x, jd.seqlen_offsets, jd.num_candidates) + ) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@torch.no_grad() +def test_low_precision_cached_and_full(dtype): + torch.manual_seed(42) + layer = TransformerInferLayer( + config(bf16=dtype == torch.bfloat16, fp16=dtype == torch.float16), 0, "cpu" + ) + x = torch.randn(9, 12).to(dtype) + jd = jagged([9], [3]) + metadata = make_metadata([6], 1, 2, 4, dtype) + got = layer.forward_naive(1, 9, x, jd, metadata) + expected = reference(layer.double(), x.double(), 3) + torch.testing.assert_close(got.double(), expected, atol=0.025, rtol=0.02) + + +def test_invalid_lengths_are_rejected(): + layer = TransformerInferLayer(config(), 0, "cpu") + with pytest.raises(RuntimeError, match="max_seq_len"): + layer(torch.randn(25, 12), offsets([25]), torch.tensor([1])) + with pytest.raises(RuntimeError, match="candidate count"): + layer(torch.randn(5, 12), offsets([5]), torch.tensor([6])) + + +def test_non_affine_input_norm_and_empty_batch_tokens(): + layer = TransformerInferLayer(config(learnable_input_layernorm=False), 0, "cpu") + x = torch.randn(5, 12) + torch.testing.assert_close( + layer(x, offsets([5]), torch.tensor([2])), reference(layer, x, 2) + ) + empty = layer(torch.empty(0, 12), offsets([0, 0]), torch.tensor([0, 0])) + assert empty.shape == (0, 12) + + +def test_matching_transformer_checkpoint_and_wrong_backbone_rejection(): + class Dense(nn.Module): + def __init__(self): + super().__init__() + self._backbone, self._use_exportable = "transformer", False + self._hstu_block = nn.Module() + self._hstu_block._attention_layers = nn.ModuleList( + [TransformerInferLayer(config(), 0, "cpu")] + ) + + original, loaded = Dense(), Dense() + state = copy.deepcopy(original.state_dict()) + state[ + "_embedding_collection._data_parallel_embedding_collection.embeddings.item.weight" + ] = torch.randn(2, 3) + load_dense_state_dict(loaded, state, strict=False) + for a, b in zip(original.parameters(), loaded.parameters()): + torch.testing.assert_close(a, b, rtol=0, atol=0) + with pytest.raises(RuntimeError, match="Checkpoint does not match transformer"): + load_dense_state_dict( + loaded, + {"_hstu_block._attention_layers.0._linear_uvqk_weight": torch.randn(1)}, + strict=False, + ) + + +def test_hstu_checkpoint_transposition_is_preserved(): + dense = nn.Module() + dense._backbone, dense._use_exportable = "hstu", False + dense._hstu_block = nn.Module() + layer = nn.Module() + layer._linear_uvqk = nn.Linear(3, 8) + layer._linear_proj = nn.Linear(2, 3, bias=False) + layer._linear_uvqk_weight = torch.empty(3, 8) + layer._linear_proj_weight = torch.empty(2, 3) + dense._hstu_block._attention_layers = nn.ModuleList([layer]) + prefix = "_hstu_block._attention_layers.0." + uvqk, proj, bias = torch.randn(3, 8), torch.randn(2, 3), torch.randn(8) + load_dense_state_dict( + dense, + { + prefix + "_linear_uvqk_weight": uvqk, + prefix + "_linear_uvqk_bias": bias, + prefix + "_linear_proj_weight": proj, + }, + ) + torch.testing.assert_close(layer._linear_uvqk.weight, uvqk.T) + torch.testing.assert_close(layer._linear_uvqk_weight, uvqk) + torch.testing.assert_close(layer._linear_proj_weight, proj) + + +def test_export_real_layer_dynamic_batch_and_tokens(tmp_path): + layer = TransformerInferLayer(config(), 0, "cpu").eval() + t = torch.export.Dim("tokens", min=2, max=40) + b = torch.export.Dim("batch", min=1, max=4) + args = (torch.randn(9, 12), offsets([4, 5]), torch.tensor([1, 2])) + ep = torch.export.export(layer, args, dynamic_shapes=({0: t}, {0: b + 1}, {0: b})) + path = tmp_path / "transformer.pt2" + torch.export.save(ep, path) + replay = torch.export.load(path).module() + for lengths, candidates in [ + ([2], [1]), + ([3, 5, 8], [1, 2, 3]), + ([0, 4, 6, 12], [0, 0, 1, 2]), + ([10, 11, 12], [3, 4, 5]), + ]: + inputs = ( + torch.randn(sum(lengths), 12), + offsets(lengths), + torch.tensor(candidates), + ) + torch.testing.assert_close(replay(*inputs), layer(*inputs)) + + +class CachedAttention(nn.Module): + """Expose the production paged reader through tensor-only export inputs.""" + + def __init__(self, layer): + super().__init__() + self._layer = layer + + def forward(self, x, off, candidates, table, ids, indptr, history): + """Run attention against an already populated cache.""" + q, k, v = self._layer.project_qkv(x) + return self._layer.finish( + x, + self._layer.attention( + q, k, v, off, candidates, table, ids, indptr, history + ), + ) + + +@dataclass +class CacheExportMetadata: + """Hold cache tensors used by the CPU export fixture.""" + + kv_cache_table: list + kv_indices: torch.Tensor + kv_indptr: torch.Tensor + total_history_lengths: torch.Tensor + kv_onload_handle: object = None + + +@dataclass +class JaggedExportMetadata: + """Hold packed sequence offsets and candidate counts for export.""" + + seqlen_offsets: torch.Tensor + num_candidates: torch.Tensor + + +class CachedLayer(nn.Module): + """Expose production cache append and attention as tensor-only inputs.""" + + def __init__(self, layer): + super().__init__() + self._layer = layer + + def forward(self, x, off, candidates, table, ids, indptr, history): + """Append new history to the input cache and compute layer outputs.""" + metadata = CacheExportMetadata( + kv_cache_table=[table], + kv_indices=ids, + kv_indptr=indptr, + total_history_lengths=history, + kv_onload_handle=None, + ) + jd = JaggedExportMetadata(seqlen_offsets=off, num_candidates=candidates) + return self._layer.forward_naive( + candidates.shape[0], x.shape[0], x, jd, metadata + ) + + +def test_export_cpu_cache_append_and_replay(tmp_path): + layer = TransformerInferLayer(config(export_mode=True), 0, "cpu").eval() + wrapper = CachedLayer(layer) + md = make_metadata([4, 7], 1, 2, 4, torch.float32) + md.kv_cache_table[0].normal_() + args = ( + torch.randn(7, 12), + offsets([3, 4]), + torch.tensor([1, 2]), + md.kv_cache_table[0], + md.kv_indices, + md.kv_indptr, + md.total_history_lengths, + ) + ep = torch.export.export(wrapper, args) + torch.export.save(ep, tmp_path / "cache_append.pt2") + replay = torch.export.load(tmp_path / "cache_append.pt2").module() + for _ in range(2): + args[0].normal_() + eager_args = tuple(x.clone() for x in args) + replay_args = tuple(x.clone() for x in args) + expected = wrapper(*eager_args) + actual = replay(*replay_args) + torch.testing.assert_close(actual, expected) + torch.testing.assert_close(replay_args[3], eager_args[3]) + + +def test_export_real_paged_reader_attention(tmp_path): + layer = TransformerInferLayer(config(), 0, "cpu").eval() + wrapper = CachedAttention(layer) + md = make_metadata([4, 7], 1, 2, 4, torch.float32) + md.kv_cache_table[0].normal_() + args = ( + torch.randn(7, 12), + offsets([3, 4]), + torch.tensor([1, 2]), + md.kv_cache_table[0], + md.kv_indices, + md.kv_indptr, + md.total_history_lengths, + ) + ep = torch.export.export(wrapper, args) + torch.export.save(ep, tmp_path / "cached.pt2") + replay = torch.export.load(tmp_path / "cached.pt2").module() + torch.testing.assert_close(replay(*args), wrapper(*args)) + # Cache content is a runtime input, not baked into the exported program. + args[3].normal_() + torch.testing.assert_close(replay(*args), wrapper(*args)) + + +@pytest.mark.skipif( + not torch.cuda.is_available(), reason="requires CUDA and compiled paged_kvcache_ops" +) +def test_cuda_append_and_graph_matches_eager(): + import paged_kvcache_ops # noqa: F401 + + layer = TransformerInferLayer( + config(hidden_size=64, head_dim=32, fp16=True), 0 + ).eval() + x = torch.randn(7, 64, device="cuda", dtype=torch.float16) + jd = jagged([3, 4], [1, 2]) + jd.seqlen_offsets, jd.num_candidates = ( + jd.seqlen_offsets.cuda(), + jd.num_candidates.cuda(), + ) + md = make_metadata([2, 2], 1, 2, 32, torch.float16) + md.kv_cache_table = [md.kv_cache_table[0].cuda()] + for attr in ("kv_indices", "kv_indptr", "total_history_lengths"): + setattr(md, attr, getattr(md, attr).cuda()) + md.batch_indices = torch.tensor([0, 0, 1, 1], dtype=torch.int32, device="cuda") + md.position = torch.tensor([0, 1, 0, 1], dtype=torch.int32, device="cuda") + md.kv_last_page_len = torch.tensor([2, 2], dtype=torch.int32, device="cuda") + md.new_history_nnz_cuda = torch.tensor([4], dtype=torch.int32, device="cuda") + with torch.inference_mode(): + expected = layer(x, jd.seqlen_offsets, jd.num_candidates) + eager = layer.forward_naive(2, 7, x, jd, md) + torch.testing.assert_close(eager, expected) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + for _ in range(3): + layer.forward_input(2, 7, x, jd, md) + layer.forward_output(2, 7, x, jd, md) + torch.cuda.current_stream().wait_stream(stream) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + layer.forward_input(2, 7, x, jd, md) + output = layer.forward_output(2, 7, x, jd, md) + graph.replay() + torch.testing.assert_close(output, expected) diff --git a/examples/hstu/training/pretrain_gr_ranking.py b/examples/hstu/training/pretrain_gr_ranking.py index f54cdc376..f0957d9ba 100644 --- a/examples/hstu/training/pretrain_gr_ranking.py +++ b/examples/hstu/training/pretrain_gr_ranking.py @@ -85,6 +85,10 @@ def main(): trainer_args.pipeline_type == "prefetch" ) network_args = NetworkArgs() + if network_args.backbone != "hstu": + raise ValueError( + "This training entry point supports HSTU only; Transformer is an inference backbone" + ) optimizer_args = OptimizerArgs() tp_args = TensorModelParallelArgs() diff --git a/examples/hstu/training/pretrain_gr_retrieval.py b/examples/hstu/training/pretrain_gr_retrieval.py index a034f2104..e22918643 100644 --- a/examples/hstu/training/pretrain_gr_retrieval.py +++ b/examples/hstu/training/pretrain_gr_retrieval.py @@ -82,6 +82,10 @@ def main(): caching=trainer_args.pipeline_type == "prefetch" ) network_args = NetworkArgs() + if network_args.backbone != "hstu": + raise ValueError( + "This training entry point supports HSTU only; Transformer is an inference backbone" + ) optimizer_args = OptimizerArgs() tp_args = TensorModelParallelArgs() diff --git a/examples/hstu/utils/gin_config_args.py b/examples/hstu/utils/gin_config_args.py index 95e0b98ce..ca272396b 100644 --- a/examples/hstu/utils/gin_config_args.py +++ b/examples/hstu/utils/gin_config_args.py @@ -383,6 +383,10 @@ class NetworkArgs: disable_contextual_mask: bool = False + # Inference backbone. Training examples continue to construct HSTU models. + backbone: str = "hstu" + transformer_ffn_dim: Optional[int] = None + def __post_init__(self): assert self.dtype_str in [ "bfloat16", @@ -390,6 +394,8 @@ def __post_init__(self): ], "Only support bfloat16 and float16 precision for Network." assert self.kernel_backend.lower() in ["cutlass", "triton", "pytorch"] + if self.backbone not in ("hstu", "transformer"): + raise ValueError(f"Unknown inference backbone: {self.backbone}") @gin.configurable