iree-org / iree-org/wave

Reductions only work on `workgroup_dim` 0

Open
#211 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
59
Forks
32
PR merge metrics
No merged PRs in 30d

Description

Reduction dim needs to be `workgroup_dim` 0 to work.

We should warn about this or make it work for the other dimensions.

```python
def get_rmsnorm_wave_rsqrt(shape, eps: float = 1e-6):
M = tkl.sym.M
N = tkl.sym.N
BLOCK_M = tkl.sym.BLOCK_M
BLOCK_N = tkl.sym.BLOCK_N
ADDRESS_SPACE = tkl.sym.ADDRESS_SPACE
threads_per_wave = 64

constraints: list[tkw.Constraint] = [
tkw.HardwareConstraint(
threads_per_wave=threads_per_wave,
vector_shapes={M: 1, N: BLOCK_N },
)
]
constraints += [tkw.WorkgroupConstraint(M, BLOCK_M, 1)] # <--- we cannot swap these !
constraints += [tkw.WorkgroupConstraint(N, BLOCK_N, 0)]
constraints += [tkw.WaveConstraint(M, BLOCK_M)]
constraints += [tkw.WaveConstraint(N, BLOCK_N)]

@tkw.wave(constraints)
def rmsnorm(
a: tkl.Memory[M, N, ADDRESS_SPACE, tkl.bf16],
weight: tkl.Memory[N, ADDRESS_SPACE, tkl.bf16],
c: tkl.Memory[M, N, ADDRESS_SPACE, tkl.bf16],
):
length_embedding = tkl.Register[M, tkl.f32](N)
eps_reg = tkl.Register[M, tkl.f32](eps)
a_reg = tkw.read(a)
a_reg = tkw.cast(a_reg, tkl.f32)
sq = a_reg * a_reg
mean = tkw.sum(sq, dim=N, block=False) / length_embedding + eps_reg
rms = tkw.sqrt(mean)
rms_broad = tkw.broadcast(rms, [M, N])
a_scaled = a_reg / rms_broad
w_reg = tkw.read(weight)
w_reg = tkw.cast(w_reg, tkl.f32)
w_broad = tkw.broadcast(w_reg, [M, N])
output = a_scaled * w_broad
output = tkw.cast(output, tkl.bf16)
tkw.write(output, c)

options = WaveCompileOptions(
subs={
M: shape[0],
N: shape[1],
BLOCK_M: 1,
BLOCK_N: shape[1],
ADDRESS_SPACE: GLOBAL_ADDRESS_SPACE,
},
canonicalize=True,
use_fast_math=True,
)
options = set_default_run_config(options)
return wave_compile(options, rmsnorm)
```

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.