Files
apache--tvm/python
Xijing Wang 82293c8c11 [Relax][Frontend][KVCache] Extend masked sequence prefill to causal left-padding (#19431)
This PR extends `_attention_sequence_prefill_with_mask` to support a
second mask regime for decoder-style embedding workloads.

### Summary

- Keep the existing right-padded bidirectional behavior as
`mask_mode="padded"`.
- Add `mask_mode="causal_padded_left"` for left-padded causal sequence
prefill.
- Add a `softmax_update_causal_padded_left` macro for the online softmax
mask.
- Add tests for causal left-padding with zero, full, mixed, and GQA
valid lengths.

### Motivation

This is a TVM-side kernel dependency for the first-class embedding
serving work tracked in mlc-ai/mlc-llm#3451.

The existing masked sequence prefill kernel supports encoder-style
batches where real tokens occupy the valid prefix `[0, valid_len)` and
padding is on the right.

Decoder-style embedding batches, such as the decoder-only embedding
path, commonly left-pad variable-length inputs so the final real token /
EOS lands at the same final column across the batch. This allows
last-token pooling to read `output[:, -1, :]`, while still requiring
causal masking within each valid suffix.

For each batch row:

- `mask_mode="padded"`: real tokens are `[0, valid_len)`.
- `mask_mode="causal_padded_left"`: real tokens are `[seq_len -
valid_len, seq_len)`, with `col <= row`.

### Testing

- `git diff --check`
- Attempted:
`python -m pytest -q
tests/python/relax/test_frontend_nn_llm_sequence_prefill_masked.py -k
'causal_padded_left or valid_len_mixed'`
2026-04-24 22:52:34 -04:00
..