strided_view → partition_view conversion reference
cutile-rs does not support make_strided_view. The IR generated by the Triton-TileIR
backend uses strided_view; it must be converted to partition_view to be
implementable in cutile-rs.
Address mapping formula
For both view kinds, load_view_tko computes the starting element position
inside the tensor_view from the index:
partition_view: element_offset[d] = index[d] × tile_shape[d]
strided_view: element_offset[d] = index[d] × traversal_strides[d]Decision rule
For each dimension d of a strided_view, let:
- T = tile_shape[d]
- S = traversal_strides[d]
- I = the set of all index values actually passed to load/store_view_tko on this dimension
Necessary and sufficient condition for convertibility: for every dimension d,
every index value i in I satisfies (i × S) % T == 0.
Converted partition index: index_partition[d] = index_strided[d] × S / T
Three cases
| Case | Condition | Conversion |
|---|---|---|
| Trivial | S == T | Drop traversal_strides; index unchanged |
| Convertible | S ≠ T, but every actual index × S is divisible by T | index_new = index_old × S / T |
| Not convertible | Some index makes (index × S) % T ≠ 0 | strided_view accesses a non-tile-aligned position that partition_view cannot express |
Common pattern: traversal_strides = [1,1,...,1]
Almost all Triton-TileIR strided_views are traversal_strides=[1,1,...,1] (from Triton-TileIR's
tl.make_block_ptr semantics). In this case S=1 and the condition simplifies to:
every actual index value must be a multiple of tile_shape[d].
Typical patterns and how to convert them:
Pattern 1: outer dimension tile_shape=1, index is arbitrary
strided_view<tile=(1, 1, 256, 128), traversal_strides=[1, 1, 1, 1], ...>
load_view_tko %sview[%batch, %head, %offset, 0]dim 0,1: T=1, S=1, any index × 1 / 1 = index. Unchanged. dim 2: T=256, S=1, offset must be a multiple of 256. dim 3: T=128, S=1, index fixed at 0. 0 × 1 / 128 = 0. Unchanged.
Conversion: find the source of offset (usually muli %blockId, %tile_size) and
extract the factor.
// Triton-TileIR original
%27 = muli %blockId_x, 256 // offset = blockId_x * 256
load_view_tko %sview[%batch, %head, %27, 0]
// after conversion
load_view_tko %pview[%batch, %head, %blockId_x, 0]Key operation: delete muli %blockId_x, 256 and use %blockId_x directly as
the partition index.
Pattern 2: cnt × tile_size inside a loop
// Triton-TileIR original: loop carries an extra cnt variable
for ... iter_values(..., %cnt = 0) {
%44 = muli %cnt, 128 // offset = cnt * TILE_N
load_view_tko %sview[%b, %h, %44, 0]
...
%next_cnt = addi %cnt, 1
continue ..., %next_cnt
}
// after conversion: use loopIdx directly, no cnt needed
for %loopIdx in (0 to hi, step 1) {
load_view_tko %pview[%b, %h, %loopIdx, 0]
}Key operations:
- Delete the loop's
cntcarry variable - Delete
muli %cnt, tile_size - Delete
addi %cnt, 1 - Use
%loopIdxin place of the computed offset
Pattern 3: dim_map transpose
strided_view does not use dim_map for transpose; it transposes with permute
after the load:
// Triton-TileIR original
strided_view<tile=(1, 1, 128, 128), traversal_strides=[1, 1, 1, 1], ...>
load_view_tko %sview[%b, %h, %offset, 0] → tile<1x1x128x128>
%2d = reshape → tile<128x128>
%transposed = permute %2d [1, 0] → tile<128x128> // transpose
mmaf %q, %transposed, ...
// after conversion: dim_map transposes at view-definition time, no permute needed
partition_view<tile=(1, 1, 128, 128), ..., dim_map=[0, 1, 3, 2]>
load_view_tko %pview[%b, %h, 0, %loopIdx] → tile<1x1x128x128>
%2d = reshape → tile<128x128>
mmaf %q, %2d, ... // already in the transposed layoutKey operations:
- Add
dim_map=[0, 1, 3, 2]on the partition_view (swap dim2 and dim3) - Swap the dim2/dim3 positions of the load index
- Delete the
permute
Note: when using dim_map, the swapped dimensions in tile_shape must also be
swapped. If Triton-TileIR has tile=(1,1,128,128) + permute [1,0], then the
partition_view can be written as tile=(1,1,128,128), dim_map=[0,1,3,2] (output
tile shape unchanged), or equivalently tile=(1,1,128,128), dim_map=[0,1,3,2]
together with the swapped index.
Full conversion procedure (Triton-TileIR → cutile-rs)
Input: make_strided_view in Triton-TileIR IR + all load/store_view_tko that reference it
Step 1: Extract traversal_strides and tile_shape
Step 2: For each load/store, collect the index expression for each dimension
Step 3: For each dimension d:
if S[d] == T[d]:
index unchanged
elif S[d] == 1:
check that every index is a multiple of T[d] (of the form x * T[d])
if yes: index_new = x (extract the multiplication factor)
if no: not convertible → keep the raw-pointer approach
else:
check (index * S[d]) % T[d] == 0
if yes: index_new = index * S[d] / T[d]
if no: not convertible
Step 4: Check whether a load is immediately followed by a permute
if yes: absorb the permute into dim_map
Step 5: Emit partition_view + the simplified index
Step 6: Delete the redundant muli/addi (offset computation) and permutepadding_value handling
Triton-TileIR's strided_view usually carries padding_value = zero. partition_view also
supports padding_value; just keep it.
If Triton-TileIR has no padding and cutile-rs's partition_view also has no padding, an out-of-bounds load returns an unspecified value (matching Triton-TileIR behavior).
Reverse conversion: partition_view → strided_view
Always possible, and trivial:
partition_view<tile=(T0, T1, ...)>
→ strided_view<tile=(T0, T1, ...), traversal_strides=[T0, T1, ...]>Index unchanged, semantics identical. partition_view is a special case of strided_view.
Worked case: Triton-TileIR strided_view → CuTile partition index
| Triton-TileIR strided_view | tile | strides | actual index | converted partition index |
|---|---|---|---|---|
| Q | (1,1,256,128) | [1,1,1,1] | [b, h, bid_x×256, 0] | [b, h, bid_x, 0] |
| QPE | (1,1,256,64) | [1,1,1,1] | [b, h, bid_x×256, 0] | [b, h, bid_x, 0] |
| K | (1,1,128,128) | [1,1,1,1] | [b, h, cnt×128, 0] | [b, h, loopIdx, 0] + dim_map=[0,1,3,2] |
| KPE | (1,1,128,64) | [1,1,1,1] | [b, 0, cnt×128, 0] | [b, 0, 0, loopIdx] + dim_map=[0,1,3,2] |
| V | (1,1,128,128) | [1,1,1,1] | [b, h, cnt×128, 0] | [b, h, loopIdx, 0] |
| Out | (1,1,256,128) | [1,1,1,1] | [b, h, bid_x×256, 0] | [b, h, bid_x, 0] |
| L | (1,1,256) | [1,1,1] | [b, h, bid_x×256] | [b, h, bid_x] |
All 7 strided_views can be losslessly converted to partition_view.