Skip to content

[refactor] move segment reductions to scatter_ops with torch_scatter names - #686

Closed
tiankongdeguiji wants to merge 2 commits into
alibaba:masterfrom
tiankongdeguiji:refactor/scatter-ops
Closed

tiankongdeguiji wants to merge 2 commits into
alibaba:masterfrom
tiankongdeguiji:refactor/scatter-ops

Conversation

@tiankongdeguiji

Copy link
Copy Markdown
Collaborator

Stacked on #684 — merge that first; this diff includes its commit until then.

Summary

jagged_segment_sum / max / logsumexp reduce rows by a per-row index (index_add_ / scatter_reduce_); they never needed a jagged layout, and their lengths argument was only read for lengths.size(0). Since #684 they are also called with indices from torch.unique on non-contiguous batches, so the jagged_ prefix and the lengths parameter 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 to segment_* would have been wrong in the other direction.

This moves them to tzrec/ops/scatter_ops.py as scatter_sum / scatter_max / scatter_logsumexp(src, index, dim_size), following torch_scatter's naming, and drops lengths. jagged_segment_ids(lengths, output_size) stays in jagged_tensors.py as 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

  • New tzrec.ops.scatter_ops_test: empty groups (sum → 0, max → 0), bf16 round-trip, index in arbitrary order, scatter_logsumexp values and gradients against per-group torch.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-commit and pyrefly check clean.

🤖 Generated with Claude Code

https://claude.ai/code/session_01WsLRKMqJQHRmUKrGy7ivdJ

tiankongdeguiji and others added 2 commits September 21, 2026 20:27
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
@tiankongdeguiji
tiankongdeguiji deleted the refactor/scatter-ops branch September 21, 2026 13:02
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.

1 participant