#ifdef Q_F32
#define Q_TYPE f32
#else
#define Q_TYPE f16
#endif

#ifdef K_F32
#define K_TYPE f32
#elif defined(K_Q4_0) || defined(K_Q8_0)
#define K_TYPE u32
#else
#define K_TYPE f16
#endif

#ifdef V_F32
#define V_TYPE f32
#elif defined(V_Q4_0) || defined(V_Q8_0)
#define V_TYPE u32
#else
#define V_TYPE f16
#endif

#ifdef DST_F32
#define DST_TYPE f32
#else
#define DST_TYPE f16
#endif

#if defined(FLASH_ATTN_SCALAR_KV) || defined(K_Q4_0) || defined(K_Q8_0)
#define K_STORAGE_TYPE K_TYPE
#else
#define K_STORAGE_TYPE vec4<K_TYPE>
#endif

#if defined(FLASH_ATTN_SCALAR_KV) || defined(V_Q4_0) || defined(V_Q8_0)
#define V_STORAGE_TYPE V_TYPE
#else
#define V_STORAGE_TYPE vec4<V_TYPE>
#endif

// Just a very small float value.
const FLOAT_MIN: f32 = -1.0e9;

struct Params {
    offset_q: u32,
    offset_k: u32,
    offset_v: u32,
    offset_mask: u32,
    offset_sinks: u32,
    offset_dst: u32,

    // shapes of Q/K/V
    n_heads: u32,
    seq_len_q: u32,
    seq_len_kv: u32,

    // strides (in elements)
    stride_q1: u32,
    stride_q2: u32,
    stride_q3: u32,
    stride_k1: u32,
    stride_k2: u32,
    stride_k3: u32,
    stride_v1: u32,
    stride_v2: u32,
    stride_v3: u32,
    stride_mask3: u32,

    // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA
    q_per_kv: u32,

    // softmax params
    scale: f32,
    max_bias: f32,
    logit_softcap: f32,
    n_head_log2: f32,
    m0: f32,
    m1: f32,

#ifdef FLASH_ATTN_VEC_SPLIT
#ifdef BLK
    blk_base: u32,
    blk_nblk0: u32,
    blk_nblk1: u32,
#endif

    tmp_data_base: u32,
    tmp_stats_base: u32,
    nwg: u32,
#endif
};

@group(0) @binding(0) var<storage, read_write> Q: array<Q_TYPE>;
@group(0) @binding(1) var<storage, read_write> K: array<K_STORAGE_TYPE>;
#ifdef KV_OVERLAP
#define V K
#define MASK_BINDING 2
#else
@group(0) @binding(2) var<storage, read_write> V: array<V_STORAGE_TYPE>;
#define MASK_BINDING 3
#endif // KV_OVERLAP

#ifdef MASK
@group(0) @binding(MASK_BINDING) var<storage, read_write> mask: array<f16>;
#define SINKS_BINDING (MASK_BINDING + 1)
#else
#define SINKS_BINDING MASK_BINDING
#endif

#ifdef SINKS
@group(0) @binding(SINKS_BINDING) var<storage, read_write> sinks: array<f32>;
#define BLK_BINDING (SINKS_BINDING + 1)
#else
#define BLK_BINDING SINKS_BINDING
#endif

#ifdef FLASH_ATTN_VEC_SPLIT
#ifdef BLK
@group(0) @binding(BLK_BINDING) var<storage, read_write> blk: array<u32>;
#define TMP_BINDING (BLK_BINDING + 1)
#else
#define TMP_BINDING BLK_BINDING
#endif

@group(0) @binding(TMP_BINDING) var<storage, read_write> tmp: array<f32>;
#define DST_BINDING (TMP_BINDING + 1)
#else
#define DST_BINDING BLK_BINDING
#endif // FLASH_ATTN_VEC_SPLIT

@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<DST_TYPE>>;

#define PARAMS_BINDING (DST_BINDING + 1)
@group(0) @binding(PARAMS_BINDING) var<uniform> params: Params;
