mirror of
https://github.com/tensorflow/tensorflow.git
synced 2026-09-28 13:23:37 +08:00
Introduce `HloOpcode::kShuffle` (`HloShuffleInstruction`), which moves the
elements of its operand around along a set of `dimensions`. A single opcode
covers the whole family of data-movement patterns (rotate, reverse, permute,
...), which keeps the semantic intent available to compiler analysis, layout
assignment, and target-specific lowerings instead of forcing an expansion into
slice and concatenate during graph construction.
The pattern to apply is selected by `xla.ShuffleMode`, a `oneof` pairing each
mode with exactly the attributes that parameterize it, so a mode cannot be
combined with another mode's attributes and adding a mode adds no fields to
`HloInstructionProto`. Switches over `mode()` are exhaustive and have no
`default`, so a new mode does not compile until every consumer has handled it.
Only `rotate` is implemented here. It takes one `shifts` entry per shuffled
dimension and rotates the elements to the left.
Specifically:
* **HLO IR & Parser**: `HloShuffleInstruction` carries `dimensions` and a
`ShuffleMode`, with text serialization and parsing. The text form is
`shuffle(x), dimensions={0,1}, mode=rotate, shifts={2,5}`, where the name of
a mode is the name of its field in `ShuffleMode`.
* **Shape Inference & Verifier**: `ShapeInference::InferShuffleShape` and the
verifier reject duplicate or out-of-range `dimensions`, a missing mode, and
mode attributes inconsistent with the mode (for `rotate`, `dimensions` and
`shifts` of unequal size).
* **XlaBuilder API**: `xla::Shuffle` constructs shuffles from symbolic
handles. The new `shuffle` utility provides `MakeRotateMode` to build a rotate
mode and `NormalizeShift` to map an arbitrary shift onto `[0, dim_size)`.
* **HLO-to-MHLO / StableHLO Translation**: `HloFunctionImporter` imports
`kShuffle`, lowering each rotated dimension sequentially into
`stablehlo.slice` and `stablehlo.concatenate`
(`concat(slice(shift:), slice(0:shift))`).
* **Pass Integration**: Register `kShuffle` across core compiler passes
including instruction fusion, layout assignment, and sharding propagation.
PiperOrigin-RevId: 987630885