[refactor] move segment reductions to scatter_ops with torch_scatter names - #686
Closed
tiankongdeguiji wants to merge 2 commits into
Closed
tiankongdeguiji wants to merge 2 commits into
tiankongdeguiji wants to merge 2 commits into
Conversation
Group list-wise losses by a session feature on any rank model, not only by the per-request candidate counts the DlrmHSTU family publishes; on those models session_name can widen a list beyond one request. jrc_loss computes its session term with segment reductions instead of a batch-by-batch mask, so it no longer scales quadratically with batch size. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01WsLRKMqJQHRmUKrGy7ivdJ
…names The reductions are keyed by a per-row index and never needed a jagged layout; the jagged_ prefix and the lengths argument promised one. Name them scatter_sum / scatter_max / scatter_logsumexp(src, index, dim_size) after torch_scatter, and keep jagged_segment_ids as the jagged bridge. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01WsLRKMqJQHRmUKrGy7ivdJ
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
jagged_segment_sum / max / logsumexpreduce rows by a per-row index (index_add_/scatter_reduce_); they never needed a jagged layout, and theirlengthsargument was only read forlengths.size(0). Since #684 they are also called with indices fromtorch.uniqueon non-contiguous batches, so thejagged_prefix and thelengthsparameter now promise a layout the functions don't need and the callers don't provide. In PyTorch,segment_*(e.g.torch.segment_reduce) means contiguous-by-lengths, so renaming tosegment_*would have been wrong in the other direction.This moves them to
tzrec/ops/scatter_ops.pyasscatter_sum / scatter_max / scatter_logsumexp(src, index, dim_size), following torch_scatter's naming, and dropslengths.jagged_segment_ids(lengths, output_size)stays injagged_tensors.pyas the jagged → index bridge. Pure rename across the five callers (jrc_loss,listwise_rank_loss,onerank_sd,task_tower, ops tests); no numerics change.Test Plan
tzrec.ops.scatter_ops_test: empty groups (sum → 0, max → 0), bf16 round-trip, index in arbitrary order,scatter_logsumexpvalues and gradients against per-grouptorch.logsumexp/ softmax, empty group →-inf.jagged_tensors_test(GPU),onerank_sd_test,dlrm_hstu_onerank_test(A10),jrc_loss_test,listwise_rank_loss_test,rank_model_test,task_tower_test— pass.pre-commitandpyrefly checkclean.🤖 Generated with Claude Code
https://claude.ai/code/session_01WsLRKMqJQHRmUKrGy7ivdJ