Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 16 additions & 2 deletions examples/hstu/configs/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from . import hstu_config, inference_config, task_config
from . import hstu_config, inference_config
from .hstu_config import (
HSTUConfig,
HSTULayerType,
Expand All @@ -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",
Expand Down
26 changes: 26 additions & 0 deletions examples/hstu/configs/inference_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
Expand All @@ -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.
Expand Down Expand Up @@ -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,
)
4 changes: 4 additions & 0 deletions examples/hstu/inference/README.md
Original file line number Diff line number Diff line change
@@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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
2 changes: 2 additions & 0 deletions examples/hstu/inference/inference_gr_ranking.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
133 changes: 133 additions & 0 deletions examples/hstu/inference/transformer.md
Original file line number Diff line number Diff line change
@@ -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.<layer_index>.`:

- `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.
2 changes: 2 additions & 0 deletions examples/hstu/inference/triton/hstu_model/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
61 changes: 61 additions & 0 deletions examples/hstu/inference_aoti/export_inference_gr_ranking.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)


Expand Down
18 changes: 15 additions & 3 deletions examples/hstu/modules/hstu_block_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
]
)
Expand All @@ -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(
Expand Down
Loading