Skip to content

[feat] add DlrmHSTUOneRank (OneRank generative ranking) - #661

Merged
tiankongdeguiji merged 10 commits into
alibaba:masterfrom
VERY-YCX:feat/hstu_multi_task
Sep 18, 2026
Merged

tiankongdeguiji merged 10 commits into
alibaba:masterfrom
VERY-YCX:feat/hstu_multi_task

Conversation

@VERY-YCX

@VERY-YCX VERY-YCX commented Sep 4, 2026 •

Copy link
Copy Markdown
Collaborator

Add DlrmHSTUOneRank: OneRank generative ranking with pluggable scorers

Summary

This PR adds DlrmHSTUOneRank (model tag 602), a OneRank-style generative ranking model built on top of the existing DlrmHSTU stack. It keeps the same input preprocessing, label decoding, point-wise losses, and metrics as DlrmHSTU, but replaces the shared FusionMTLTower head with per-task scoring over per-candidate, per-task representations produced by a structured token expansion inside HSTU.

Key changes

  • Structured tokenization: each candidate is expanded into K interleaved (candidate copy, task token) pairs. The resulting task-private attention mask is encoded as two contiguous column intervals, so it fits the existing NFUNC=3 arbitrary-mask path without modifying the attention kernel.
  • Situation Discernment (optional): per-task contextual query + softmax MHCA pooling over the request's own candidates. Disabled by default in our recommended config (see Deviations below).
  • Cross-task head (optional): cascade-masked cross-task attention with strategic gradient detachment.
  • List-wise InfoNCE loss: per-request softmax over the candidate list, jagged by candidate.sequence_length. No cross-rank gather is required.
  • Pluggable scorer: dot_product (paper default), bilinear (identity-init per-task W_k), and mlp (per-task two-layer MLP). task_bias_init allows initializing per-task logits to logit(CTR_k).

Deviations from the OneRank paper

The implementation deliberately keeps SD and the cross-task head as config-gated optional modules, but our experiments on long-video TV recommendation found that:

  • the rank-1 dot_product scorer tends to plateau at the global CTR prediction early in training;
  • replacing SD + MHCA with simple mean-pooling and using the mlp scorer gives significantly better results in this scenario.

Therefore the recommended config for this use case is:

  • situation_discernment unset,
  • cross_task_head unset,
  • scorer_type = ONERANK_SCORER_MLP,
  • task_bias_init set to per-task empirical logit(CTR).

The other options are retained for ablation studies and for users who want to reproduce the original paper setup.

Backward compatibility

Only three upstream files receive subclass hooks (_build_output_modules, _build_stu_layer, _build_attn_func). The single behavioural change in stu.py is guarded by uses_arbitrary_mask and is False for all existing DlrmHSTU / UltraHSTU configs, so there is no impact on existing models.

Requirements

  • kernel: CUTLASS or PYTORCH; TRITON is not supported because the NFUNC arbitrary-mask path is not implemented there.
  • bf16/fp16 mixed precision is recommended for CUTLASS.
  • Mid-stack attention truncation and KV-cache serving are not supported for this model.

Testing

  • Unit tests for OneRankListwiseLoss covering numerical stability, masking edge cases, and gradient correctness.
  • Model-level tests for DlrmHSTUOneRank on both PYTORCH and CUTLASS kernels, including kernel consistency checks and gradient reachability for task-token parameters.

@CLAassistant

CLAassistant commented Sep 4, 2026 •

Copy link
Copy Markdown

CLA assistant check
All committers have signed the CLA.

@VERY-YCX VERY-YCX changed the title Feat/hstu_multi_task [feat] add DlrmHSTUOneRank (OneRank generative ranking) Sep 4, 2026
@tiankongdeguiji tiankongdeguiji added the claude-review Let Claude Review label Sep 5, 2026
@github-actions github-actions Bot removed the claude-review Let Claude Review label Sep 5, 2026
Comment thread tzrec/modules/gr/onerank_tokenizer.py
Comment thread tzrec/modules/gr/onerank_jagged.py Outdated
Comment thread tzrec/models/dlrm_hstu_onerank.py Outdated
Comment thread tzrec/models/dlrm_hstu_onerank_test.py
Comment thread tzrec/protos/models/multi_task_rank.proto
Comment thread tzrec/modules/gr/onerank_head.py Outdated
Comment thread tzrec/modules/gr/onerank_sd.py
@github-actions

github-actions Bot commented Sep 5, 2026

Copy link
Copy Markdown
Contributor

Code review summary

Overall a well-structured addition: the structured-tokenization design is clearly motivated, the numerical handling in the list-wise loss is careful (per-segment max shift, exact clamps, masked-out requests, temperature clamp), and the tests are strong in places — the func-tensor vs. reference-mask test and the loss value+gradient reference checks are exactly the right assertions. The stu.py refactor was verified to be a no-op for existing SLA users (SLA func intervals are subsets of the fixed causal/target/contextual masks), and the model-level flip/length bookkeeping matches the DlrmHSTU baseline.

Main issue (see inline on onerank_tokenizer.py:230): the model as wired cannot be constructed with a default config — the stu.contextual_seq_len sentinel (-1) resolves to input_preprocessor.contextual_seq_len(), which is ≥ 1 for any contextual group, and OneRankSTULayer rejects anything but 0. This fires for the PR's own test configs too (they never set stu.contextual_seq_len: 0), so either the transducer needs to pin the value or the required config needs documenting and testing. The accompanying "baseline runs with contextual_seq_len = 0" claim is also inaccurate (the baseline gets a bidirectional contextual block).

Other noteworthy items (details inline):

  • dlrm_hstu_onerank.py:249 — global-average-loss rescale uses total requests while the loss averages over valid requests only; biased cross-rank mean when mask ratios differ.
  • onerank_jagged.py:39 — ~25 hidden device→host syncs per step from repeat_interleave without output_size=; all call sites have the total statically.
  • dlrm_hstu_onerank_test.py — scorer_type variants / task_bias_init, cross-task gradient-detachment semantics, and listwise loss-input alignment have no coverage.
  • onerank_sd.py:205 — per-task SD loop and cross-task pairwise projection are batchable/project-then-expandable.
  • multi_task_rank.proto:209 — item_embedding_hidden_dim is silently ignored by this model; a few comment gaps noted there.
  • Docs: ONERANK_OVERVIEW.md is referenced twice but does not exist in the repo, and there is no docs/source/models page / generative.rst entry / README row for the new model, unlike every comparable model (dlrm_hstu, ultra_hstu, hstu_match, sid_*).

Minor notes not worth inline comments: the comment in _build_output_modules claiming config_to_kwargs "does not translate enums or repeated fields" is inaccurate (json_format.MessageToDict translates enums to name strings and repeated fields to lists, and config_to_kwargs(onerank.cross_task_head) relies on exactly that for mask_type/hybrid_chain_task_names); nan_to_num(per_request, nan=0.0) in the loss does not absorb ±Inf (defaults rewrite them to the dtype's max finite value, which a valid request would carry into the sum); and incompatible HSTU options (attention truncation, KV-cache serving) fail at first forward rather than at construction, where the analogous constraints are already validated.

Note: the code-quality leg of this review was still running when this summary was posted; follow-up comments may follow.

@github-actions

github-actions Bot commented Sep 5, 2026

Copy link
Copy Markdown
Contributor

Follow-up: the code-quality leg of the review has now finished. It independently confirmed the two headline findings already posted (the stu.contextual_seq_len construction blocker at onerank_tokenizer.py:230, and the missing ONERANK_OVERVIEW.md / user-facing docs) and found no additional issues. It also verified — so maintainers need not re-litigate — that the NFUNC func-tensor intervals match the paper layout (including no-prefix and zero-candidate edges), the task-token slot selection in _compose_output is consistent with the tokenizer ordering, the double flip in predict restores both request and within-request candidate order, and the expand-based replicas sum gradients back into the candidate correctly. No further comments from this review.

VERY-YCX pushed a commit to VERY-YCX/TorchEasyRec that referenced this pull request Sep 7, 2026
Code-review fixes, one per bot comment:

1. pin contextual_seq_len to 0 and correct the sentinel error/docstring
   (the STU group layout owns the context/target split; the flag was a
   miswired no-op that could never be anything but 0).
2. pass output_size to the 5 repeat_interleave call sites so the jagged
   expansions stay FX-traceable without a graph break.
3. unify the listwise-loss denominator to the number of requests, so
   the per-rank rescale of enable_global_average_loss is unbiased under
   DDP regardless of per-rank request counts.
4. add the missing test coverage: scorer variants (dot/bilinear/mlp),
   task_bias_init, cross-task detachment semantics, listwise reference
   resync, plus 17 head-level and 4 model-level tests.
5. close the 4 proto comment gaps (BCE/BFL condition, autotune note,
   sentinel doc, item_embedding_hidden_dim IGNORED).
6. add the twice-referenced ONERANK_OVERVIEW.md (architecture, losses,
   plateau dynamics, deltas vs the paper) and the docs model page.
7. project before the (B,K,K,D) cross-task expansion and batch the SD
   scoring -- numerically identical (verified against the loop form),
   Kx less projection GEMM.

Also fixes test_cutlass_backward_reaches_the_task_tokens: wrapping
TrainWrapper in a manual autocast never worked because the wrapper
always enters its own autocast context (enabled=False when
mixed_precision is None), which disabled the outer one and handed the
fp16/bf16-only CUTLASS kernel fp32 q/k/v. Mixed precision now goes
through the production knob (TrainWrapper mixed_precision='BF16',
driven by train_config.mixed_precision), and the test passes.

Verified: onerank_listwise_loss_test (11), onerank_head_test (17),
onerank_tokenizer_test (31), dlrm_hstu_onerank_test (13, GPU, incl.
CUTLASS forward+backward) -- all green.
@VERY-YCX VERY-YCX added the claude-review Let Claude Review label Sep 7, 2026
@github-actions github-actions Bot removed the claude-review Let Claude Review label Sep 7, 2026
Comment thread docs/source/models/dlrm_hstu_onerank.md
Comment thread docs/source/models/dlrm_hstu_onerank.md
Comment thread docs/source/models/dlrm_hstu_onerank.md Outdated
Comment thread tzrec/protos/model.proto Outdated
Comment thread tzrec/modules/gr/onerank_cross_task.py Outdated
Comment thread tzrec/models/dlrm_hstu_onerank_test.py
Comment thread tzrec/modules/gr/onerank_head_test.py Outdated
Comment thread tzrec/modules/gr/onerank_head_test.py Outdated
Comment thread tzrec/models/dlrm_hstu_onerank_test.py
Comment thread tzrec/models/dlrm_hstu_onerank.py Outdated
Comment thread tzrec/models/dlrm_hstu_onerank.py Outdated
@github-actions

github-actions Bot commented Sep 7, 2026

Copy link
Copy Markdown
Contributor

Code review — DlrmHSTUOneRank (automated, 5-area panel)

Overall this is a high-quality, well-hardened contribution: the NFUNC=3 mask-interval math checks out row-by-row (prefix/replica/task-token), the stu.py edits are verifiably a no-op for existing DlrmHSTU/UltraHSTU/SLA configs, the listwise InfoNCE is numerically robust (exact per-segment max-shift, denom >= 1 by construction, clamped temperature, all-masked batches tested finite), and the test suite goes well past the repo norm (CUTLASS↔PyTorch kernel consistency, gradient reachability into task tokens, detachment semantics).

11 inline comments posted. The ones I'd treat as blocking-or-explain:

  1. The Chinese doc's example config cannot parse — task_configs uses non-existent fields (label_names/loss_type/metrics_set), and kernel: CUTLASS + input_dropout_ratio are nested inside stu {} where neither is a field. Since ModelConfig.kernel defaults to PYTORCH, a naive fix would silently run the reference kernel against the doc's own CUTLASS requirement.
  2. model.proto oneof numbering — dlrm_hstu_onerank = 602 lands in the SID-generation block; multi-task rank models live in the 200s (dlrm_hstu = 205, ultra_hstu = 207). Oneof numbers are permanent; 208 is cheap now, expensive after release.
  3. Cross-task Strategic Gradient Detachment — detaching after _k_proj/_v_proj means those weights only ever learn from diagonal entries, contradicting both the docstring ("projection weights still learn from every pair") and the "gradients are identical" comment. Forward values are identical, so no test catches it.
  4. max_num_candidates is an unchecked data contract — a tail request exceeding it silently mis-autotunes every jagged kernel (eager) or fails mid-run inside a Triton op (compile). The tokenizer already pays the needed D→H sync; the check is free.
  5. Test sampling gaps — sequence_timestamp_is_ascending=False may never be drawn in non-nightly CI (max_examples // 5, derandomized, 8 sampled axes), and the baseline needed a deterministic ordering test for exactly this contract; OneRank's listwise regrouping makes an order slip shape-invisible. SD has no value-level test; build_cross_task_mask has none at all.

Minor items (no inline comment each)

  • pre-commit appears not run: README table isn't mdformat-normalized; onerank_head.py:215-218's magic trailing comma will be exploded by ruff-format; non-ASCII dashes in new comments outside tzrec/ops/.
  • Comment inaccuracies: onerank_tokenizer.py:137-139 ("Negative on prefix rows" — torch.remainder takes the divisor's sign, so slot is always non-negative); dlrm_hstu_onerank.py:173-175 (config_to_kwargs does translate enums/repeated fields — see inline); proto says scorer_hidden_dim is "ONERANK_SCORER_MLP only" but validation is unconditional; task_bias_init's exact-length-K rule is undocumented while the doc example shows a single value for a multi-task model.
  • ONERANK_OVERVIEW.md placement: an English user-facing explainer at repo root, outside the Sphinx tree and unreachable from any toctree, largely restating module docstrings; repo convention is Chinese user docs in docs/, implementation contracts beside the code. Consider folding section 5's plateau guidance into docs/source/models/dlrm_hstu_onerank.md and moving or dropping the root file (the proto and onerank_head.py cite it).
  • Internal artifacts in permanent surfaces: the duration-bitmask enumeration {0,1,3,7,15,31}, "p99 = 10 on the target dataset", and "~100/900+ steps" run numbers in proto comments and docstrings don't generalize — state the model-level rationale instead.
  • Loss edge cases: a zero-request batch on a rank gives 0/0 NaN that poisons all ranks through the DDP average — max(lengths.size(0), 1) completes the otherwise-thorough degenerate-path hardening (inherited from the point-wise path, but this module hardens everything else). Consider logits.float() under bf16 autocast, matching BCEWithLogitsLoss's fp32 promotion. logit_scale's gradient value is never compared against the reference (a .detach() on the scale would pass current tests).
  • Fail-late config checks: TRITON kernel and attn_truncation_* are knowable at build time but only raise on first forward; the transducer already rejects five other misconfigs at construction.
  • Perf micro-items: the stride-0 func-tensor view is re-.contiguous()'d per layer inside cutlass_hstu_mha (pre-existing for SLA, amplified by the inflated total_q); loss() adds one more scalar all-reduce over a value derivable from static shapes; _get_label re-decodes per listwise task; _preprocess returns a sum-of-maxes max_seq_len bound, looser than the max-of-sums the two D→H syncs were justified to avoid.
  • Test consistency: pure-CPU classes tagged @mark_ci_scope("h20", "gpu") while sibling CPU classes are untagged; assertTrue(compute_metric()) only asserts a non-empty dict; OneRankHSTUTransducer's construction-time guards (interleaving, return_full_embeddings, contextual sentinel resolution) are untested; unused task_weight/output_dropout_ratio knobs in the test builder; _build_stu_layer/_preprocess overrides lack docstrings.

Areas reviewed: code quality, performance, test coverage, documentation accuracy, safety/validation — all five completed.

VERY-YCX added a commit to VERY-YCX/TorchEasyRec that referenced this pull request Sep 8, 2026
Code-review fixes, one per bot comment:

1. pin contextual_seq_len to 0 and correct the sentinel error/docstring
   (the STU group layout owns the context/target split; the flag was a
   miswired no-op that could never be anything but 0).
2. pass output_size to the 5 repeat_interleave call sites so the jagged
   expansions stay FX-traceable without a graph break.
3. unify the listwise-loss denominator to the number of requests, so
   the per-rank rescale of enable_global_average_loss is unbiased under
   DDP regardless of per-rank request counts.
4. add the missing test coverage: scorer variants (dot/bilinear/mlp),
   task_bias_init, cross-task detachment semantics, listwise reference
   resync, plus 17 head-level and 4 model-level tests.
5. close the 4 proto comment gaps (BCE/BFL condition, autotune note,
   sentinel doc, item_embedding_hidden_dim IGNORED).
6. add the twice-referenced ONERANK_OVERVIEW.md (architecture, losses,
   plateau dynamics, deltas vs the paper) and the docs model page.
7. project before the (B,K,K,D) cross-task expansion and batch the SD
   scoring -- numerically identical (verified against the loop form),
   Kx less projection GEMM.

Also fixes test_cutlass_backward_reaches_the_task_tokens: wrapping
TrainWrapper in a manual autocast never worked because the wrapper
always enters its own autocast context (enabled=False when
mixed_precision is None), which disabled the outer one and handed the
fp16/bf16-only CUTLASS kernel fp32 q/k/v. Mixed precision now goes
through the production knob (TrainWrapper mixed_precision='BF16',
driven by train_config.mixed_precision), and the test passes.

Verified: onerank_listwise_loss_test (11), onerank_head_test (17),
onerank_tokenizer_test (31), dlrm_hstu_onerank_test (13, GPU, incl.
CUTLASS forward+backward) -- all green.
@VERY-YCX
VERY-YCX force-pushed the feat/hstu_multi_task branch from 0053969 to a8639c9 Compare September 8, 2026 08:41
Comment thread tzrec/models/dlrm_hstu_onerank_test.py
Comment thread tzrec/modules/gr/onerank_sd.py
Comment thread tzrec/models/dlrm_hstu_onerank.py Outdated
Comment thread tzrec/modules/gr/onerank_cross_task.py
Comment thread tzrec/modules/gr/stu.py
@github-actions

github-actions Bot commented Sep 9, 2026

Copy link
Copy Markdown
Contributor

Code review - DlrmHSTUOneRank (fresh 5-area panel: code quality, performance, test coverage, docs accuracy, security/validation)

All five areas completed against the current HEAD. This is an unusually careful PR, and it shows: we independently re-verified the pieces most likely to break silently and they hold up.

Verified sound (no action needed)

  • stu.py hooks: the non-SLA path passes byte-identical kernel arguments, and the "no-op for SLA" claim checks out - build_sla_func_tensor folds contextual_seq_len (via effective_k2) and target isolation (via H_b) into the func intervals, so the func-visible set is a subset of both fixed masks and the previously-ANDed masks were redundant. UltraHSTU is untouched by the refactor.
  • NFUNC=3 mask math row-by-row (prefix/replica/task-token intervals, no-prefix and zero-candidate edges, int32 plumbing), tokenizer slot order vs. _compose_output slot-1 extraction vs. task_configs order end-to-end.
  • Flip bookkeeping in predict()/loss() (pre-flip num_targets, post-flip head inputs, flip-back of mt_preds) is order-consistent with jagged labels, and the enable_global_average_loss rescale (N_r / avg over a total-count denominator) is an unbiased global mean after DDP averaging.
  • Loss hardening (detached per-segment max shift, exact denom.clamp(min=1), valid mask by multiply, nan_to_num, max(B,1)) prevents NaN/inf leakage into the all-reduce; the only collective is config-gated, so no rank-divergent participation.
  • Registration surface: oneof 208 continues the 200s block, which_msg/RegisterABCMeta resolves the class, README/generative.rst/proto field docs are mutually consistent, and the arXiv link is the right paper.

New inline comments from this pass

  • docs/.../dlrm_hstu_onerank.md:162 - mixed_precision: BF16 must be quoted (string field); the recommended CUTLASS example otherwise fails text-format parsing.
  • docs/.../dlrm_hstu_onerank.md:179 - model export is listed as rejected, but nothing enforces it: the cand_seq_pk guard at export_util.py:147 covers only dlrm_hstu/ultra_hstu, and the export path traces the full forward without ever hitting cached_forward.
  • onerank_jagged.py:71-75 - index_add_ accumulates in the input dtype; under the documented BF16 the recommended no-SD mean-pool builds z_k from bf16 atomics.
  • onerank_tokenizer.py:551-556 - output postprocessor runs over the full 2K-row expansion before half the rows are dropped, and the non-contiguous slice pins the expanded buffer through backward.
  • onerank_tokenizer.py:489-490 - both per-step D-to-H syncs are derivable from ints the preprocessor already synced (max_new_targets == max_targets * group_size exactly).
  • dlrm_hstu_onerank_test.py:766,803 - the two CUTLASS tests (the only end-to-end NFUNC verification) inherit only the class-level "gpu" scope; repo precedent marks CUTLASS-wheel tests ("h20", "gpu").
  • dlrm_hstu_onerank_test.py:395 - hypothesis GPU test lacks the sibling-test teardown_example/cleanup_cuda_memory().
  • stu.py:536-547 - the SLA+CUTLASS behavioural no-op is untested through STULayer; a parity test would be cheap insurance.
  • Two echo prior-panel comments and reinforce rather than add: SD per-task channel/stack order still has no value-level test (onerank_sd.py:211-220), and the class docstring still ties bf16/fp16 to both kernels (dlrm_hstu_onerank.py:69).

Prior-panel items we confirmed fixed at this HEAD: oneof 602 -> 208; docs moved into docs/source/models/ with rst/README entries; example-config field names; cross-task detachment now projects the detached input (pinned by gradient-equality tests); max_num_candidates runtime bound; output_size= on every repeat_interleave; zero-batch max(B,1) guard; logit_scale gradient compared against the reference; deterministic num_targets-order test.

Prior-panel items still open at this HEAD (no new inline; see the existing threads): temperature_init < 0.01 starts above the ln(100) clamp, so the learnable scale gets zero gradient forever (onerank_listwise_loss.py:67); the interleaving guard probes a private _enable_interleaving attribute fail-open (onerank_tokenizer.py:414); attn_truncation_* is the one incompatible knob rejected at first forward instead of construction (onerank_tokenizer.py:320); predict() duplicates the order-critical flip/assembly skeleton of the baseline instead of hooking it (dlrm_hstu_onerank.py:278); the example comment claiming FusionMTLTower requires mlp (md:101, same claim in dlrm_hstu_onerank_test.py:183).

Minor notes (no inline)

  • The PR description says the recommended config leaves SD and cross_task_head unset, but the shipped doc example sets both - worth reconciling one way or the other.
  • FP16 is documented for CUTLASS but only BF16 is tested; the listwise loss under low precision is asserted isfinite only (float64 parity is tested at unit level).
  • New parameterized.expand tables omit name_func=parameterized_name_func (AGENTS.md); the loss test expands to opaque _0/_1/_2.
  • Per-step aside already acknowledged in-code: the prefix split-expand-concat round trip copies the (long) prefix twice, and the prefix-timestamp concat is dead work since _compose_output discards the left half - architectural, fine to defer.

@VERY-YCX
VERY-YCX force-pushed the feat/hstu_multi_task branch from ba1e328 to e31698d Compare September 10, 2026 03:38
@VERY-YCX

Copy link
Copy Markdown
Collaborator Author

Housekeeping update: the 8 development commits (feature + four review-fix rounds + two follow-ups) have been squashed into the single head commit e31698d so the PR now reads as one complete change against master. The tree is byte-identical to the previously reviewed state plus two new fixes:

  • 48bf972 — quote mixed_precision in the docs (it is a string field; the bare BF16 failed text_format.Parse)
  • 39d056b — make the listwise-loss denominator FX-traceable (torch.fx.wrap leaf), fixing a TraceError under torchrec's train-pipeline rewrite that unit tests never exercise

All review feedback that was already addressed by the fix rounds (4f62fc9 / 90da9b5) has been replied to and resolved inline (22 threads). The remaining 17 open threads are the medium/low items from the last review round (doc-accuracy wording, minor test coverage, and optional perf suggestions) — they are tracked and will be handled in follow-up commits on this same branch.

VERY-YCX added a commit to VERY-YCX/TorchEasyRec that referenced this pull request Sep 11, 2026
…cision

Review threads 1/2 of PR alibaba#661.

- DlrmHSTU.predict delegates the middle scoring block to a
  _score_targets hook, so the pre-flip num_targets capture, both
  fx_flip_tensor_dict flips, the _output_to_prediction_impl loop and
  the TARGET_REPEAT_INTERLEAVE_KEY publication are single-sourced in
  the base class. The default hook is the old block moved verbatim
  (DlrmHSTU / UltraHSTU behavior unchanged); DlrmHSTUOneRank drops its
  ~50-line predict() copy and overrides only the hook, and the
  module-local fx_flip_tensor_dict fx-wrap goes away with the call
  site.
- jagged_segment_sum accumulates through
  promote_types(values.dtype, fp32) and casts back: bf16/fp16 inputs
  pool through an fp32 accumulator (index_add_ is not on autocast's
  promote list, unlike sum), fp32 is bit-identical and fp64 keeps full
  precision. promote_types keeps the dtype choice a graph node instead
  of control flow on values.dtype, which fx symbolic tracing rejects
  when it inlines this body.

Verified: dlrm_hstu / dlrm_hstu_onerank / ultra_hstu / listwise loss /
onerank head / tokenizer suites 99/99 incl. fx-trace and AOT-inductor
graph variants; ruff check+format clean; bit-level dtype matrix on
CPU+CUDA (fp32/fp64 equal to the previous formula, bf16/fp16 CUDA now
accumulate in fp32).
@VERY-YCX
VERY-YCX marked this pull request as ready for review September 11, 2026 06:27
Comment thread tzrec/protos/models/multi_task_rank.proto Outdated
Comment thread tzrec/modules/gr/onerank_tokenizer.py Outdated
Comment thread tzrec/modules/gr/onerank_sd.py Outdated
Comment thread tzrec/models/dlrm_hstu_onerank.py
Comment thread tzrec/models/dlrm_hstu_onerank.py Outdated
Comment thread tzrec/models/dlrm_hstu.py Outdated
Comment thread tzrec/modules/gr/onerank_head.py Outdated
Comment thread tzrec/modules/gr/onerank_jagged.py Outdated
Comment thread tzrec/modules/gr/onerank_head.py Outdated
Add DlrmHSTUOneRank (model tag 602): a OneRank-style generative ranking
model on top of the DlrmHSTU stack. It keeps DlrmHSTU's input
preprocessing, label decoding, point-wise losses and metrics, and
replaces the shared FusionMTLTower head with per-task scoring over
per-candidate, per-task representations:

- structured tokenization: each candidate expands into K interleaved
  (candidate copy, task token) pairs; the task-private attention mask
  is encoded as two contiguous column intervals on the existing
  NFUNC=3 arbitrary-mask path (kernel: CUTLASS/PYTORCH, not TRITON)
- optional situation discernment (SD): per-task request-context query
  over the candidate pool, with cross-task attention under a cascade
  mask (duration tasks may read the click task, not vice versa)
- pluggable scorers: DOT (paper's rank-1 inner product), BILINEAR
  (identity-init W), MLP; plus per-task bias init (logit of global CTR)
- optional listwise InfoNCE loss per task (same-request candidates as
  negatives), with an fx-traceable denominator compatible with
  torchrec's train_pipeline rewrite
- losses/labels/metrics inherited from fusion_mtl_tower task_configs;
  existing models are untouched (guarded by uses_arbitrary_mask)

Includes unit tests (mask semantics, scorer shapes, listwise loss),
the dlrm_hstu_onerank.md user guide (config examples verified with
text_format.Parse, including the quoted mixed_precision string), and
dataset/pipeline docs.
…evel tests

Docs (dlrm_hstu_onerank.md, multi_task_rank.proto):
- same-config A/B: note that the baseline resolves an unset
  stu.contextual_seq_len to the contextual feature count (a bidirectional
  attention block) while OneRank pins it to 0, so that deviation is part
  of the measured metric delta.
- fusion_mtl_tower.mlp is optional for OneRank (the tower is never
  built): drop the "required" wording; keep it consistent with the
  "ignored fields" note and the proto comment.
- state the actual listwise validity rule (each request needs >= 1
  positive AND >= 1 negative) instead of the ambiguous "positive
  coverage close to 100%"; the OneRankListwiseLoss proto comment
  carried the same phrasing and now states the rule too.
- the unsupported-configs table now says where the KV-cache
  NotImplementedError actually fires (serving-side cached_forward, not
  at export), and that tzrec.export neither rejects OneRank nor
  requires the cand_seq_pk key its siblings need.

Docstrings (dlrm_hstu_onerank.py):
- scope the fp16/bf16 requirement to the CUTLASS kernel; the PYTORCH
  reference kernel is fp32-capable.

Construction-time guards (onerank_tokenizer.py, preprocessors.py):
- reject attn_truncation_split_layer / attn_truncation_tail_len at
  construction: a reused dlrm_hstu block with truncation tuning used to
  only fail via OneRankSTULayer.truncate_input's NotImplementedError on
  the first training step, after torchrun and sharding were fully up.
- replace the cross-module private-attribute probe
  (_enable_interleaving) with a public has_interleaving() accessor on
  InputPreprocessor; the base class returns False, so unknown
  preprocessor types fail closed.

Tests:
- value-level pin of the listwise data wiring: a bare
  OneRankListwiseLoss evaluated on hand-decoded bitmask labels and the
  published logits/num_targets (both timestamp directions), plus a
  request-permutation oracle that catches a regression of the mt_preds
  flip-back under descending timestamps; both were mutation-verified
  (raw labels / removed flip-back each fail the new tests).
- JaggedCrossAttention against a dense per-segment softmax reference
  (num_heads=2, random projections) and a direct jagged_softmax test
  with non-uniform logits.
- OneRankSituationDiscernment per-task wiring: an identity/uniform
  test pinning out[:, k] == mean(task channel k) and a module-pairing
  reference test over the per-task lists.
- method-level mark_ci_scope("h20", "gpu") on both end-to-end CUTLASS
  tests: the gpu lane lacks the fbgemm_gpu_hstu wheel, so without the
  h20 tag they would run on no per-PR lane at all.
- teardown_example -> cleanup_cuda_memory() between hypothesis
  examples, matching the sibling HSTU test files.
…cision

Review threads 1/2 of PR alibaba#661.

- DlrmHSTU.predict delegates the middle scoring block to a
  _score_targets hook, so the pre-flip num_targets capture, both
  fx_flip_tensor_dict flips, the _output_to_prediction_impl loop and
  the TARGET_REPEAT_INTERLEAVE_KEY publication are single-sourced in
  the base class. The default hook is the old block moved verbatim
  (DlrmHSTU / UltraHSTU behavior unchanged); DlrmHSTUOneRank drops its
  ~50-line predict() copy and overrides only the hook, and the
  module-local fx_flip_tensor_dict fx-wrap goes away with the call
  site.
- jagged_segment_sum accumulates through
  promote_types(values.dtype, fp32) and casts back: bf16/fp16 inputs
  pool through an fp32 accumulator (index_add_ is not on autocast's
  promote list, unlike sum), fp32 is bit-identical and fp64 keeps full
  precision. promote_types keeps the dtype choice a graph node instead
  of control flow on values.dtype, which fx symbolic tracing rejects
  when it inlines this body.

Verified: dlrm_hstu / dlrm_hstu_onerank / ultra_hstu / listwise loss /
onerank head / tokenizer suites 99/99 incl. fx-trace and AOT-inductor
graph variants; ruff check+format clean; bit-level dtype matrix on
CPU+CUDA (fp32/fp64 equal to the previous formula, bf16/fp16 CUDA now
accumulate in fp32).
W1 [6]/[4]/[2b]/[9]: _predict_impl hook renames, scorer_type as a
plain proto string, build_onerank_func_tensor moved to
ops/hstu_attention_utils.py, OneRankPredictionHead moved to
task_tower.py with its tests re-homed (4-way split).

W2 [1]/[3]/[8]: max_num_candidates deleted -- max_seq_len now bounds
the *inflated* sequence and the tokenizer guards it at runtime; SD
switched to varlen_attn with device dispatch as an fx leaf;
onerank_jagged dissolved into private helpers and deleted.

W3 [5]: listwise loss generalized to LossConfig.listwise_rank_loss,
consumed by the RankModel base class (_init_loss_impl/_loss_impl/
_output_to_prediction_impl branches); fx_avg_batch_size promoted to
fx_util; onerank model slimmed (init_loss/loss overrides removed);
loss module renamed listwise_rank_loss.py.

Also adds tzrec/utils/fx_util_test.py pinning the AVG semantics of
fx_avg_batch_size's distributed branch.
…aba#670)

Review [2a]: now that alibaba#670 landed HSTU_ARBITRARY_NFUNC = 5, candidate
replicas are retired and each candidate expands into exactly K + 1
tokens, as the paper writes it.

- build_onerank_func_tensor: 5-wide three-interval encoding.  t_k's
  visible set [0, H_b) u {e^C_i} u {t_k} splits naturally into the
  (NFUNC + 1) // 2 = 3 column intervals because the mutually invisible
  t_1..t_{k-1} sit between e^C_i and t_k -- the very reason replicas
  existed under NFUNC=3.
- OneRankTokenizer: group_size is K + 1; forward concatenates the
  candidate with the task tokens (torch.cat, no replica expand).
- Inflation factors follow across tests and docs: default max_seq_len
  148 -> 132, runtime-rejection case 26 -> 18, docs example 2328 ->
  2208.
- group_size validation drops the even-parity rule (odd sizes are
  legal now); the floor is 2.
- Tests rewritten against the paper layout: reference mask, slot
  indices, stamp positions; the replica-gradient test is replaced by
  one covering the single candidate row plus the per-group task-token
  gradient aggregation.
- CUTLASS equivalence tests now require the fn5 wheel (CI h20 lane);
  locally verified through the PyTorch reference kernel
  (_decode_attn_func_to_mask) instead.
@VERY-YCX
VERY-YCX force-pushed the feat/hstu_multi_task branch from 26348ad to a76065b Compare September 15, 2026 10:59
Comment thread tzrec/models/rank_model.py
Comment thread tzrec/modules/gr/onerank_cross_task.py Outdated
Comment thread tzrec/modules/task_tower_test.py Outdated
Comment thread tzrec/modules/task_tower_test.py Outdated
Comment thread tzrec/modules/task_tower_test.py Outdated
Comment thread tzrec/loss/listwise_rank_loss.py Outdated
- copyright year 2025 -> 2026 across all PR-added files (they were
  created in 2026)
- task_tower_test: move the scorer constants, _inputs and _head into
  OneRankPredictionHeadTest as class attributes / classmethods; the
  parameterized.expand cases are built with an explicit loop because
  comprehension iterables cannot see the class scope
- rank_model: listwise_rank_loss publishes its own logits/probs in
  _output_to_prediction_impl, making it a valid standalone objective;
  drop the cross-loss dependency and its loud-failure branch, flip the
  pinning test to a standalone-trains-fine one, and refresh the docs
  and ListwiseRankLoss docstring accordingly
The explicit class-body loop building _SCORER_CASES needed a 'del' of its
loop variables, which pyrefly flags as potentially uninitialized
(unbound-name) since the loop body may not run. Replace it with a
comprehension over itertools.product: the single 'for' iterable is
evaluated in the class scope (nested comprehension iterables would not
see it), and the loop variables stay comprehension-local, so there is
nothing to delete. Case names and ordering are unchanged.
listwise_rank_loss.py, onerank_sd.py, and task_tower.py each carried a
byte-identical private copy of _jagged_segment_ids/_sum/_max. Promote
them to public jagged_segment_ids/jagged_segment_sum/jagged_segment_max
in tzrec/ops/jagged_tensors.py (pure PyTorch compositions, so no Kernel
dispatch) and import them at all call sites. Value-level tests for the
shared contracts (empty-segment handling, reduced-precision fp32
accumulation round-trip) live in jagged_tensors_test.py.
@tiankongdeguiji
tiankongdeguiji merged commit bac9ac6 into alibaba:master Sep 18, 2026
7 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants