#!/usr/bin/env python3 """FlashInfer 0.6.14 prefill/decode API probe with a PyTorch reference.""" from __future__ import annotations import json import torch from flashinfer.decode import single_decode_with_kv_cache from flashinfer.prefill import single_prefill_with_kv_cache def main() -> None: if not torch.cuda.is_available(): raise SystemExit("CUDA GPU required") torch.manual_seed(7) qo_len, kv_len, heads, head_dim = 8, 16, 8, 128 q = torch.randn(qo_len, heads, head_dim, device="cuda", dtype=torch.float16) k = torch.randn(kv_len, heads, head_dim, device="cuda", dtype=torch.float16) v = torch.randn_like(k) prefill = single_prefill_with_kv_cache(q, k, v, causal=False, kv_layout="NHD") decode = single_decode_with_kv_cache(q[-1], k, v, kv_layout="NHD") print( json.dumps( { "prefill_shape": list(prefill.shape), "decode_shape": list(decode.shape), "dtype": str(prefill.dtype), "receipt_scope": "synthetic_attention_shapes_not_glm_5_2", }, indent=2, ) ) if __name__ == "__main__": main()