[feat] add DlrmHSTUOneRank (OneRank generative ranking) - #661
Conversation
Code review summaryOverall 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 Main issue (see inline on Other noteworthy items (details inline):
Minor notes not worth inline comments: the comment in Note: the code-quality leg of this review was still running when this summary was posted; follow-up comments may follow. |
|
Follow-up: the code-quality leg of the review has now finished. It independently confirmed the two headline findings already posted (the |
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.
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 11 inline comments posted. The ones I'd treat as blocking-or-explain:
Minor items (no inline comment each)
Areas reviewed: code quality, performance, test coverage, documentation accuracy, safety/validation — all five completed. |
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.
0053969 to
a8639c9
Compare
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)
New inline comments from this pass
Prior-panel items we confirmed fixed at this HEAD: oneof 602 -> 208; docs moved into Prior-panel items still open at this HEAD (no new inline; see the existing threads): Minor notes (no inline)
|
ba1e328 to
e31698d
Compare
|
Housekeeping update: the 8 development commits (feature + four review-fix rounds + two follow-ups) have been squashed into the single head commit
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. |
…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).
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.
26348ad to
a76065b
Compare
- 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.
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 existingDlrmHSTUstack. It keeps the same input preprocessing, label decoding, point-wise losses, and metrics asDlrmHSTU, but replaces the sharedFusionMTLTowerhead with per-task scoring over per-candidate, per-task representations produced by a structured token expansion inside HSTU.Key changes
Kinterleaved(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.candidate.sequence_length. No cross-rank gather is required.dot_product(paper default),bilinear(identity-init per-taskW_k), andmlp(per-task two-layer MLP).task_bias_initallows initializing per-task logits tologit(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:
dot_productscorer tends to plateau at the global CTR prediction early in training;mlpscorer gives significantly better results in this scenario.Therefore the recommended config for this use case is:
situation_discernmentunset,cross_task_headunset,scorer_type = ONERANK_SCORER_MLP,task_bias_initset 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 instu.pyis guarded byuses_arbitrary_maskand isFalsefor all existingDlrmHSTU/UltraHSTUconfigs, so there is no impact on existing models.Requirements
kernel: CUTLASSorPYTORCH;TRITONis not supported because the NFUNC arbitrary-mask path is not implemented there.Testing
OneRankListwiseLosscovering numerical stability, masking edge cases, and gradient correctness.DlrmHSTUOneRankon bothPYTORCHandCUTLASSkernels, including kernel consistency checks and gradient reachability for task-token parameters.