diff --git a/mamba_ssm/ops/triton/mamba3/angle_dt.py b/mamba_ssm/ops/triton/mamba3/angle_dt.py index 5363efaa2..acfa85966 100644 --- a/mamba_ssm/ops/triton/mamba3/angle_dt.py +++ b/mamba_ssm/ops/triton/mamba3/angle_dt.py @@ -194,22 +194,23 @@ def angle_dt_fwd( else: grid = (nheads, batch) - angle_dt_fwd_kernel[grid]( - out, output_state, - angle, dt, init_state, cu_seqlens, - out.stride(0), out.stride(1), out.stride(2), out.stride(3), - stride_output_state[0], stride_output_state[1], stride_output_state[2], - angle.stride(0), angle.stride(1), angle.stride(2), angle.stride(3), - dt.stride(0), dt.stride(1), dt.stride(2), - stride_init[0], stride_init[1], stride_init[2], - stride_cu_seqlen, - seqlen, dim, - CHUNK_SIZE=chunk_size, - BLOCK_D=BLOCK_D, - HAS_INIT_STATE=HAS_INIT_STATE, - RETURN_OUTPUT_STATE=return_output_state, - IS_VARLEN=is_varlen, - ) + with torch.cuda.device(out.device.index): + angle_dt_fwd_kernel[grid]( + out, output_state, + angle, dt, init_state, cu_seqlens, + out.stride(0), out.stride(1), out.stride(2), out.stride(3), + stride_output_state[0], stride_output_state[1], stride_output_state[2], + angle.stride(0), angle.stride(1), angle.stride(2), angle.stride(3), + dt.stride(0), dt.stride(1), dt.stride(2), + stride_init[0], stride_init[1], stride_init[2], + stride_cu_seqlen, + seqlen, dim, + CHUNK_SIZE=chunk_size, + BLOCK_D=BLOCK_D, + HAS_INIT_STATE=HAS_INIT_STATE, + RETURN_OUTPUT_STATE=return_output_state, + IS_VARLEN=is_varlen, + ) if return_output_state: return out, output_state @@ -411,22 +412,23 @@ def angle_dt_bwd( else: grid = (nheads, batch) - angle_dt_bwd_kernel[grid]( - grad_angle, grad_dt, grad_init_state if has_init_state else grad_init_dummy, - grad_out, grad_output_state, angle, dt, cu_seqlens, - grad_angle.stride(0), grad_angle.stride(1), grad_angle.stride(2), grad_angle.stride(3), - grad_dt.stride(0), grad_dt.stride(1), grad_dt.stride(2), - stride_grad_init[0], stride_grad_init[1], stride_grad_init[2], - grad_out.stride(0), grad_out.stride(1), grad_out.stride(2), grad_out.stride(3), - stride_grad_output_state[0], stride_grad_output_state[1], stride_grad_output_state[2], - angle.stride(0), angle.stride(1), angle.stride(2), angle.stride(3), - dt.stride(0), dt.stride(1), dt.stride(2), - stride_cu_seqlen, - seqlen, dim, - CHUNK_SIZE=chunk_size, - BLOCK_D=BLOCK_D, - HAS_INIT_STATE=has_init_state, - HAS_GRAD_OUTPUT_STATE=HAS_GRAD_OUTPUT_STATE, - IS_VARLEN=is_varlen, - ) + with torch.cuda.device(grad_angle.device.index): + angle_dt_bwd_kernel[grid]( + grad_angle, grad_dt, grad_init_state if has_init_state else grad_init_dummy, + grad_out, grad_output_state, angle, dt, cu_seqlens, + grad_angle.stride(0), grad_angle.stride(1), grad_angle.stride(2), grad_angle.stride(3), + grad_dt.stride(0), grad_dt.stride(1), grad_dt.stride(2), + stride_grad_init[0], stride_grad_init[1], stride_grad_init[2], + grad_out.stride(0), grad_out.stride(1), grad_out.stride(2), grad_out.stride(3), + stride_grad_output_state[0], stride_grad_output_state[1], stride_grad_output_state[2], + angle.stride(0), angle.stride(1), angle.stride(2), angle.stride(3), + dt.stride(0), dt.stride(1), dt.stride(2), + stride_cu_seqlen, + seqlen, dim, + CHUNK_SIZE=chunk_size, + BLOCK_D=BLOCK_D, + HAS_INIT_STATE=has_init_state, + HAS_GRAD_OUTPUT_STATE=HAS_GRAD_OUTPUT_STATE, + IS_VARLEN=is_varlen, + ) return grad_angle, grad_dt, grad_init_state \ No newline at end of file diff --git a/mamba_ssm/ops/triton/mamba3/grouped_head_reduction.py b/mamba_ssm/ops/triton/mamba3/grouped_head_reduction.py index c91e891df..08f27be4f 100644 --- a/mamba_ssm/ops/triton/mamba3/grouped_head_reduction.py +++ b/mamba_ssm/ops/triton/mamba3/grouped_head_reduction.py @@ -189,25 +189,26 @@ def reduce_grouped_qk_grads_and_bias_triton( n_blocks = triton.cdiv(N, block_n) grid = (num_row_blocks, R, num_qk_groups * n_blocks) - _reduce_grouped_qk_grads_and_bias_partial_kernel[grid]( - dq_raw, - dk_raw, - dq_out, - dk_out, - dq_bias_partial, - dk_bias_partial, - total_rows, - R, - H, - num_qk_groups, - N, - group_size, - n_blocks, - block_m, - block_n, - store_grouped, - num_warps=8, - ) + with torch.cuda.device(dq_raw.device.index): + _reduce_grouped_qk_grads_and_bias_partial_kernel[grid]( + dq_raw, + dk_raw, + dq_out, + dk_out, + dq_bias_partial, + dk_bias_partial, + total_rows, + R, + H, + num_qk_groups, + N, + group_size, + n_blocks, + block_m, + block_n, + store_grouped, + num_warps=8, + ) dq_bias = torch.empty((H, R, N), dtype=torch.float32, device=dq_raw.device) dk_bias = torch.empty_like(dq_bias) diff --git a/mamba_ssm/ops/triton/mamba3/mamba3_mimo_utils.py b/mamba_ssm/ops/triton/mamba3/mamba3_mimo_utils.py index 5de93aea7..e8ff7d8a6 100644 --- a/mamba_ssm/ops/triton/mamba3/mamba3_mimo_utils.py +++ b/mamba_ssm/ops/triton/mamba3/mamba3_mimo_utils.py @@ -593,18 +593,19 @@ def bwd_dtrap_ddt_triton( dtrap = torch.zeros_like(trap) grid = (B, H, nchunks) - bwd_dtrap_ddt_kernel[grid]( - trap, dt, dfactor, dgamma_diag, - ddt, dtrap, - trap.stride(0), trap.stride(1), trap.stride(2), - dt.stride(0), dt.stride(1), dt.stride(2), - dfactor.stride(0), dfactor.stride(1), dfactor.stride(2), - dgamma_diag.stride(0), dgamma_diag.stride(1), dgamma_diag.stride(2), - ddt.stride(0), ddt.stride(1), ddt.stride(2), - dtrap.stride(0), dtrap.stride(1), dtrap.stride(2), - S, - chunk_size, - ) + with torch.cuda.device(trap.device.index): + bwd_dtrap_ddt_kernel[grid]( + trap, dt, dfactor, dgamma_diag, + ddt, dtrap, + trap.stride(0), trap.stride(1), trap.stride(2), + dt.stride(0), dt.stride(1), dt.stride(2), + dfactor.stride(0), dfactor.stride(1), dfactor.stride(2), + dgamma_diag.stride(0), dgamma_diag.stride(1), dgamma_diag.stride(2), + ddt.stride(0), ddt.stride(1), ddt.stride(2), + dtrap.stride(0), dtrap.stride(1), dtrap.stride(2), + S, + chunk_size, + ) return ddt, dtrap def compute_dacs_segsum_triton_varlen( @@ -672,18 +673,19 @@ def compute_dacs_segsum_triton_varlen( state_chunk_in_seq[chunk_start:chunk_end] = torch.arange(n, dtype=torch.int32, device=da.device) grid = (B, H, nchunks) - dacs_segsum_kernel_varlen[grid]( - da, da_cs, da_cs_rev, segsum, cu_seqlens, state_seq_mapping, state_chunk_in_seq, - da.stride(0), da.stride(1), da.stride(2), - da_cs.stride(0), da_cs.stride(1), da_cs.stride(2), - da_cs_rev.stride(0), da_cs_rev.stride(1), da_cs_rev.stride(2), - segsum.stride(0), segsum.stride(1), segsum.stride(2), - segsum.stride(3), segsum.stride(4), - cu_seqlens.stride(0), state_seq_mapping.stride(0), state_chunk_in_seq.stride(0), - S, - num_sequences, - chunk_size, - ) + with torch.cuda.device(da.device.index): + dacs_segsum_kernel_varlen[grid]( + da, da_cs, da_cs_rev, segsum, cu_seqlens, state_seq_mapping, state_chunk_in_seq, + da.stride(0), da.stride(1), da.stride(2), + da_cs.stride(0), da_cs.stride(1), da_cs.stride(2), + da_cs_rev.stride(0), da_cs_rev.stride(1), da_cs_rev.stride(2), + segsum.stride(0), segsum.stride(1), segsum.stride(2), + segsum.stride(3), segsum.stride(4), + cu_seqlens.stride(0), state_seq_mapping.stride(0), state_chunk_in_seq.stride(0), + S, + num_sequences, + chunk_size, + ) return da_cs, da_cs_rev, segsum @@ -708,16 +710,17 @@ def compute_dacs_segsum_triton( segsum = torch.empty(B, H, nchunks, chunk_size, chunk_size, device=da.device, dtype=da.dtype) grid = (B, H, nchunks) - dacs_segsum_kernel[grid]( - da, da_cs, da_cs_rev, segsum, - da.stride(0), da.stride(1), da.stride(2), - da_cs.stride(0), da_cs.stride(1), da_cs.stride(2), - da_cs_rev.stride(0), da_cs_rev.stride(1), da_cs_rev.stride(2), - segsum.stride(0), segsum.stride(1), segsum.stride(2), - segsum.stride(3), segsum.stride(4), - S, - chunk_size, - ) + with torch.cuda.device(da.device.index): + dacs_segsum_kernel[grid]( + da, da_cs, da_cs_rev, segsum, + da.stride(0), da.stride(1), da.stride(2), + da_cs.stride(0), da_cs.stride(1), da_cs.stride(2), + da_cs_rev.stride(0), da_cs_rev.stride(1), da_cs_rev.stride(2), + segsum.stride(0), segsum.stride(1), segsum.stride(2), + segsum.stride(3), segsum.stride(4), + S, + chunk_size, + ) return da_cs, da_cs_rev, segsum @@ -1336,31 +1339,32 @@ def bwd_dadt_fused_triton_varlen( dadt_out = torch.zeros(B, H, S, device=ddA_cs.device, dtype=torch.float32) grid = (B, H, nchunks_global) - bwd_dadt_cumsum_fused_kernel_varlen[grid]( - ddA_cs, ddA_cs_rev, dA_cs, dA_cs_rev, dadt_out, - cu_seqlens, state_seq_mapping, state_chunk_in_seq, - ddA_cs.stride(0), ddA_cs.stride(1), ddA_cs.stride(2), - cu_seqlens.stride(0), - state_seq_mapping.stride(0), - state_chunk_in_seq.stride(0), - S=S, - CHUNK_SIZE=chunk_size, - ) + with torch.cuda.device(dSSdA.device.index): + bwd_dadt_cumsum_fused_kernel_varlen[grid]( + ddA_cs, ddA_cs_rev, dA_cs, dA_cs_rev, dadt_out, + cu_seqlens, state_seq_mapping, state_chunk_in_seq, + ddA_cs.stride(0), ddA_cs.stride(1), ddA_cs.stride(2), + cu_seqlens.stride(0), + state_seq_mapping.stride(0), + state_chunk_in_seq.stride(0), + S=S, + CHUNK_SIZE=chunk_size, + ) - bwd_segsum_dadt_kernel_varlen[grid]( - dSSdA, SSdA, dadt_out, - cu_seqlens, state_seq_mapping, state_chunk_in_seq, - dSSdA.stride(0), dSSdA.stride(1), dSSdA.stride(2), - dSSdA.stride(3), dSSdA.stride(4), - SSdA.stride(0), SSdA.stride(1), SSdA.stride(2), - SSdA.stride(3), SSdA.stride(4), - dadt_out.stride(0), dadt_out.stride(1), dadt_out.stride(2), - cu_seqlens.stride(0), - state_seq_mapping.stride(0), - state_chunk_in_seq.stride(0), - S=S, - CHUNK_SIZE=chunk_size, - ) + bwd_segsum_dadt_kernel_varlen[grid]( + dSSdA, SSdA, dadt_out, + cu_seqlens, state_seq_mapping, state_chunk_in_seq, + dSSdA.stride(0), dSSdA.stride(1), dSSdA.stride(2), + dSSdA.stride(3), dSSdA.stride(4), + SSdA.stride(0), SSdA.stride(1), SSdA.stride(2), + SSdA.stride(3), SSdA.stride(4), + dadt_out.stride(0), dadt_out.stride(1), dadt_out.stride(2), + cu_seqlens.stride(0), + state_seq_mapping.stride(0), + state_chunk_in_seq.stride(0), + S=S, + CHUNK_SIZE=chunk_size, + ) return dadt_out @@ -1388,21 +1392,22 @@ def bwd_dtrap_ddt_triton_varlen( dtrap = torch.zeros_like(trap) grid = (B, H, nchunks_global) - bwd_dtrap_ddt_kernel_varlen[grid]( - trap, dt, dfactor, dgamma_diag, ddt, dtrap, - cu_seqlens, state_seq_mapping, state_chunk_in_seq, - trap.stride(0), trap.stride(1), trap.stride(2), - dt.stride(0), dt.stride(1), dt.stride(2), - dfactor.stride(0), dfactor.stride(1), dfactor.stride(2), - dgamma_diag.stride(0), dgamma_diag.stride(1), dgamma_diag.stride(2), - ddt.stride(0), ddt.stride(1), ddt.stride(2), - dtrap.stride(0), dtrap.stride(1), dtrap.stride(2), - cu_seqlens.stride(0), - state_seq_mapping.stride(0), - state_chunk_in_seq.stride(0), - S=S, - CHUNK_SIZE=chunk_size, - ) + with torch.cuda.device(trap.device.index): + bwd_dtrap_ddt_kernel_varlen[grid]( + trap, dt, dfactor, dgamma_diag, ddt, dtrap, + cu_seqlens, state_seq_mapping, state_chunk_in_seq, + trap.stride(0), trap.stride(1), trap.stride(2), + dt.stride(0), dt.stride(1), dt.stride(2), + dfactor.stride(0), dfactor.stride(1), dfactor.stride(2), + dgamma_diag.stride(0), dgamma_diag.stride(1), dgamma_diag.stride(2), + ddt.stride(0), ddt.stride(1), ddt.stride(2), + dtrap.stride(0), dtrap.stride(1), dtrap.stride(2), + cu_seqlens.stride(0), + state_seq_mapping.stride(0), + state_chunk_in_seq.stride(0), + S=S, + CHUNK_SIZE=chunk_size, + ) return ddt, dtrap diff --git a/mamba_ssm/ops/triton/mamba3/mamba3_siso_bwd.py b/mamba_ssm/ops/triton/mamba3/mamba3_siso_bwd.py index a8a0c17e6..08d9d097f 100644 --- a/mamba_ssm/ops/triton/mamba3/mamba3_siso_bwd.py +++ b/mamba_ssm/ops/triton/mamba3/mamba3_siso_bwd.py @@ -165,24 +165,25 @@ def compute_dzdo( def grid(META): return (triton.cdiv(seqlen, META["CHUNK_SIZE"]), nheads, batch) - mamba3_siso_bwd_kernel_dzdo[grid]( - do, z, o, - dz, do_scaled, - # DO strides - do.stride(0), do.stride(1), do.stride(2), do.stride(3), - # Z strides - z.stride(0), z.stride(1), z.stride(2), z.stride(3), - # O strides - o.stride(0), o.stride(1), o.stride(2), o.stride(3), - # Dz strides - dz.stride(0), dz.stride(1), dz.stride(2), dz.stride(3), - # DO_scaled strides - do_scaled.stride(0), do_scaled.stride(1), do_scaled.stride(2), do_scaled.stride(3), - # Dimensions - seqlen, headdim_v, - # Compile-time constants - HEADDIM_V=HEADDIM_V, - ) + with torch.cuda.device(do.device.index): + mamba3_siso_bwd_kernel_dzdo[grid]( + do, z, o, + dz, do_scaled, + # DO strides + do.stride(0), do.stride(1), do.stride(2), do.stride(3), + # Z strides + z.stride(0), z.stride(1), z.stride(2), z.stride(3), + # O strides + o.stride(0), o.stride(1), o.stride(2), o.stride(3), + # Dz strides + dz.stride(0), dz.stride(1), dz.stride(2), dz.stride(3), + # DO_scaled strides + do_scaled.stride(0), do_scaled.stride(1), do_scaled.stride(2), do_scaled.stride(3), + # Dimensions + seqlen, headdim_v, + # Compile-time constants + HEADDIM_V=HEADDIM_V, + ) return dz, do_scaled @@ -733,64 +734,65 @@ def compute_dqkv( grid = (nheads, batch) # Launch kernel - mamba3_siso_bwd_kernel_dqkv[grid]( - q, k, v, da_cs, da_cs_sum, qk_dot, D, SSM_States, do, d_ossm_state, Cu_Seqlens, - dq, dk, dv, dAdt, dQK, dD, d_issm_state, - # Q strides - q.stride(0), q.stride(1), q.stride(2), q.stride(3), - # K strides - k.stride(0), k.stride(1), k.stride(2), k.stride(3), - # V strides - v.stride(0), v.stride(1), v.stride(2), v.stride(3), - # DA_CS strides - da_cs.stride(0), da_cs.stride(1), da_cs.stride(2), - # DA_CS_SUM strides - da_cs_sum.stride(0), da_cs_sum.stride(1), da_cs_sum.stride(2), - # QK_Dot strides - qk_dot.stride(0), qk_dot.stride(1), qk_dot.stride(2), - # D stride - D.stride(0) if D is not None else 0, - # SSM_States strides: (batch, nheads, headdim_v, nchunks*headdim_qk) - SSM_States.stride(0), SSM_States.stride(1), SSM_States.stride(2), - SSM_States.stride(3), - # dO strides - do.stride(0), do.stride(1), do.stride(2), do.stride(3), - # d_ossm_state strides - d_ossm_state.stride(0) if d_ossm_state is not None else 0, - d_ossm_state.stride(1) if d_ossm_state is not None else 0, - d_ossm_state.stride(2) if d_ossm_state is not None else 0, - d_ossm_state.stride(3) if d_ossm_state is not None else 0, - # Cu_Seqlens strides - Cu_Seqlens.stride(0) if Cu_Seqlens is not None else 0, - # dQ strides - dq.stride(0), dq.stride(1), dq.stride(2), dq.stride(3), - # dK strides - dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3), - # dV strides - dv.stride(0), dv.stride(1), dv.stride(2), dv.stride(3), - # dAdt strides - dAdt.stride(0), dAdt.stride(1), dAdt.stride(2), - # dQK strides - dQK.stride(0), dQK.stride(1), dQK.stride(2), - # dD strides - dD.stride(0) if D is not None else 0, - dD.stride(1) if D is not None else 0, - # d_issm_state strides - d_issm_state.stride(0) if d_issm_state is not None else 0, - d_issm_state.stride(1) if d_issm_state is not None else 0, - d_issm_state.stride(2) if d_issm_state is not None else 0, - d_issm_state.stride(3) if d_issm_state is not None else 0, - # Dimensions - seqlen, nheads_qk, headdim_qk, headdim_v, - # Compile-time constants - CHUNK_SIZE=chunk_size, - HEADDIM_QK=HEADDIM_QK, - HEADDIM_V=HEADDIM_V, - RECOMPUTE_MASK=False, - HAS_D_OSSM_STATE=d_ossm_state is not None, - RETURN_D_ISSM_STATE=has_input_state, - IS_VARLEN=is_varlen, - ) + with torch.cuda.device(q.device.index): + mamba3_siso_bwd_kernel_dqkv[grid]( + q, k, v, da_cs, da_cs_sum, qk_dot, D, SSM_States, do, d_ossm_state, Cu_Seqlens, + dq, dk, dv, dAdt, dQK, dD, d_issm_state, + # Q strides + q.stride(0), q.stride(1), q.stride(2), q.stride(3), + # K strides + k.stride(0), k.stride(1), k.stride(2), k.stride(3), + # V strides + v.stride(0), v.stride(1), v.stride(2), v.stride(3), + # DA_CS strides + da_cs.stride(0), da_cs.stride(1), da_cs.stride(2), + # DA_CS_SUM strides + da_cs_sum.stride(0), da_cs_sum.stride(1), da_cs_sum.stride(2), + # QK_Dot strides + qk_dot.stride(0), qk_dot.stride(1), qk_dot.stride(2), + # D stride + D.stride(0) if D is not None else 0, + # SSM_States strides: (batch, nheads, headdim_v, nchunks*headdim_qk) + SSM_States.stride(0), SSM_States.stride(1), SSM_States.stride(2), + SSM_States.stride(3), + # dO strides + do.stride(0), do.stride(1), do.stride(2), do.stride(3), + # d_ossm_state strides + d_ossm_state.stride(0) if d_ossm_state is not None else 0, + d_ossm_state.stride(1) if d_ossm_state is not None else 0, + d_ossm_state.stride(2) if d_ossm_state is not None else 0, + d_ossm_state.stride(3) if d_ossm_state is not None else 0, + # Cu_Seqlens strides + Cu_Seqlens.stride(0) if Cu_Seqlens is not None else 0, + # dQ strides + dq.stride(0), dq.stride(1), dq.stride(2), dq.stride(3), + # dK strides + dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3), + # dV strides + dv.stride(0), dv.stride(1), dv.stride(2), dv.stride(3), + # dAdt strides + dAdt.stride(0), dAdt.stride(1), dAdt.stride(2), + # dQK strides + dQK.stride(0), dQK.stride(1), dQK.stride(2), + # dD strides + dD.stride(0) if D is not None else 0, + dD.stride(1) if D is not None else 0, + # d_issm_state strides + d_issm_state.stride(0) if d_issm_state is not None else 0, + d_issm_state.stride(1) if d_issm_state is not None else 0, + d_issm_state.stride(2) if d_issm_state is not None else 0, + d_issm_state.stride(3) if d_issm_state is not None else 0, + # Dimensions + seqlen, nheads_qk, headdim_qk, headdim_v, + # Compile-time constants + CHUNK_SIZE=chunk_size, + HEADDIM_QK=HEADDIM_QK, + HEADDIM_V=HEADDIM_V, + RECOMPUTE_MASK=False, + HAS_D_OSSM_STATE=d_ossm_state is not None, + RETURN_D_ISSM_STATE=has_input_state, + IS_VARLEN=is_varlen, + ) # Add output V state gradients to the last token if d_ov_state is not None: @@ -1279,55 +1281,56 @@ def compute_dqktheta( # Grid: (nchunks, batch) grid = (nchunks, batch) - mamba3_siso_bwd_kernel_rotary_bias_angles[grid]( - # Input tensors - q, k, scale, gamma, q_bias, k_bias, angles, dq_in, dk_in, dqk, - # Output tensors - dq, dk, dangles, dscale, dgamma, dq_bias_partial, dk_bias_partial, - # Q strides: (batch, seqlen, nheads_qk, headdim_qk) - q.stride(0), q.stride(1), q.stride(2), q.stride(3), - # K strides - k.stride(0), k.stride(1), k.stride(2), k.stride(3), - # Scale strides: (batch, nheads, seqlen) - scale.stride(0), scale.stride(1), scale.stride(2), - # SGamma strides - gamma.stride(0), gamma.stride(1), gamma.stride(2), - # Q_bias strides: (nheads, headdim_qk) - q_bias.stride(0), q_bias.stride(1), - # K_bias strides - k_bias.stride(0), k_bias.stride(1), - # Angles strides: (batch, seqlen, nheads, headdim_qk//2) - angles.stride(0), angles.stride(1), angles.stride(2), angles.stride(3), - # dQ_in strides: (batch, seqlen, nheads, headdim_qk) - dq_in.stride(0), dq_in.stride(1), dq_in.stride(2), dq_in.stride(3), - # dK_in strides - dk_in.stride(0), dk_in.stride(1), dk_in.stride(2), dk_in.stride(3), - # dQK strides: (batch, nheads, seqlen) - dqk.stride(0), dqk.stride(1), dqk.stride(2), - # Output tensors - # dQ strides: (batch, seqlen, nheads_qk, headdim_qk) - dq.stride(0), dq.stride(1), dq.stride(2), dq.stride(3), - # dK strides - dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3), - # dAngles strides: (batch, seqlen, nheads, headdim_qk//2) - dangles.stride(0), dangles.stride(1), dangles.stride(2), dangles.stride(3), - # dScale strides: (batch, nheads, seqlen) - dscale.stride(0), dscale.stride(1), dscale.stride(2), dscale.stride(3), - # dSGamma strides - dgamma.stride(0), dgamma.stride(1), dgamma.stride(2), dgamma.stride(3), - # dQ_bias_partial strides: (batch, nchunks, nheads, headdim_qk) - dq_bias_partial.stride(0), dq_bias_partial.stride(1), - dq_bias_partial.stride(2), dq_bias_partial.stride(3), - # dK_bias_partial strides - dk_bias_partial.stride(0), dk_bias_partial.stride(1), - dk_bias_partial.stride(2), dk_bias_partial.stride(3), - # Sizes - seqlen, nheads_qk, nheads, headdim_qk, headdim_angles, - CHUNK_SIZE=chunk_size, - HEADDIM_QK=HEADDIM_QK, - BLOCK_HEADDIM_QK=BLOCK_HEADDIM_QK, - GQA_RATIO=GQA_RATIO, - ) + with torch.cuda.device(q.device.index): + mamba3_siso_bwd_kernel_rotary_bias_angles[grid]( + # Input tensors + q, k, scale, gamma, q_bias, k_bias, angles, dq_in, dk_in, dqk, + # Output tensors + dq, dk, dangles, dscale, dgamma, dq_bias_partial, dk_bias_partial, + # Q strides: (batch, seqlen, nheads_qk, headdim_qk) + q.stride(0), q.stride(1), q.stride(2), q.stride(3), + # K strides + k.stride(0), k.stride(1), k.stride(2), k.stride(3), + # Scale strides: (batch, nheads, seqlen) + scale.stride(0), scale.stride(1), scale.stride(2), + # SGamma strides + gamma.stride(0), gamma.stride(1), gamma.stride(2), + # Q_bias strides: (nheads, headdim_qk) + q_bias.stride(0), q_bias.stride(1), + # K_bias strides + k_bias.stride(0), k_bias.stride(1), + # Angles strides: (batch, seqlen, nheads, headdim_qk//2) + angles.stride(0), angles.stride(1), angles.stride(2), angles.stride(3), + # dQ_in strides: (batch, seqlen, nheads, headdim_qk) + dq_in.stride(0), dq_in.stride(1), dq_in.stride(2), dq_in.stride(3), + # dK_in strides + dk_in.stride(0), dk_in.stride(1), dk_in.stride(2), dk_in.stride(3), + # dQK strides: (batch, nheads, seqlen) + dqk.stride(0), dqk.stride(1), dqk.stride(2), + # Output tensors + # dQ strides: (batch, seqlen, nheads_qk, headdim_qk) + dq.stride(0), dq.stride(1), dq.stride(2), dq.stride(3), + # dK strides + dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3), + # dAngles strides: (batch, seqlen, nheads, headdim_qk//2) + dangles.stride(0), dangles.stride(1), dangles.stride(2), dangles.stride(3), + # dScale strides: (batch, nheads, seqlen) + dscale.stride(0), dscale.stride(1), dscale.stride(2), dscale.stride(3), + # dSGamma strides + dgamma.stride(0), dgamma.stride(1), dgamma.stride(2), dgamma.stride(3), + # dQ_bias_partial strides: (batch, nchunks, nheads, headdim_qk) + dq_bias_partial.stride(0), dq_bias_partial.stride(1), + dq_bias_partial.stride(2), dq_bias_partial.stride(3), + # dK_bias_partial strides + dk_bias_partial.stride(0), dk_bias_partial.stride(1), + dk_bias_partial.stride(2), dk_bias_partial.stride(3), + # Sizes + seqlen, nheads_qk, nheads, headdim_qk, headdim_angles, + CHUNK_SIZE=chunk_size, + HEADDIM_QK=HEADDIM_QK, + BLOCK_HEADDIM_QK=BLOCK_HEADDIM_QK, + GQA_RATIO=GQA_RATIO, + ) # Reshape outputs back to original layout dscale = torch.sum(dscale, dim=2) # Sum over headdim blocks @@ -1382,35 +1385,36 @@ def apply_dk_state_post( grid = (nheads, num_sequences) - mamba3_siso_bwd_kernel_dk_state_post[grid]( - # Input tensors - d_ok_state, angles, k, k_bias, Cu_Seqlens, - # Output tensors - dk, dk_bias, dangles, - # dK_State strides: (batch, nheads, headdim_qk) - d_ok_state.stride(0), d_ok_state.stride(1), d_ok_state.stride(2), - # Angles strides: (batch, seqlen, nheads, headdim_angles) - angles.stride(0), angles.stride(1), angles.stride(2), angles.stride(3), - # K strides: (batch, seqlen, nheads_qk, headdim_qk) - k.stride(0), k.stride(1), k.stride(2), k.stride(3), - # K_bias strides: (nheads, headdim_qk) - k_bias.stride(0), k_bias.stride(1), - # Cu_Seqlens strides: (num_sequences + 1,) - Cu_Seqlens.stride(0) if is_varlen else 0, - # dK strides: (batch, seqlen, nheads_qk, headdim_qk) - dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3), - # dK_bias strides: (nheads, headdim_qk) - dk_bias.stride(0), dk_bias.stride(1), - # dAngles strides: (batch, seqlen, nheads, headdim_angles) - dangles.stride(0), dangles.stride(1), dangles.stride(2), dangles.stride(3), - # Dimensions - seqlen, headdim_qk, headdim_angles, - HEADDIM_QK=HEADDIM_QK, - GQA_RATIO=GQA_RATIO, - IS_VARLEN=is_varlen, - num_warps=2, - num_stages=3, - ) + with torch.cuda.device(d_ok_state.device.index): + mamba3_siso_bwd_kernel_dk_state_post[grid]( + # Input tensors + d_ok_state, angles, k, k_bias, Cu_Seqlens, + # Output tensors + dk, dk_bias, dangles, + # dK_State strides: (batch, nheads, headdim_qk) + d_ok_state.stride(0), d_ok_state.stride(1), d_ok_state.stride(2), + # Angles strides: (batch, seqlen, nheads, headdim_angles) + angles.stride(0), angles.stride(1), angles.stride(2), angles.stride(3), + # K strides: (batch, seqlen, nheads_qk, headdim_qk) + k.stride(0), k.stride(1), k.stride(2), k.stride(3), + # K_bias strides: (nheads, headdim_qk) + k_bias.stride(0), k_bias.stride(1), + # Cu_Seqlens strides: (num_sequences + 1,) + Cu_Seqlens.stride(0) if is_varlen else 0, + # dK strides: (batch, seqlen, nheads_qk, headdim_qk) + dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3), + # dK_bias strides: (nheads, headdim_qk) + dk_bias.stride(0), dk_bias.stride(1), + # dAngles strides: (batch, seqlen, nheads, headdim_angles) + dangles.stride(0), dangles.stride(1), dangles.stride(2), dangles.stride(3), + # Dimensions + seqlen, headdim_qk, headdim_angles, + HEADDIM_QK=HEADDIM_QK, + GQA_RATIO=GQA_RATIO, + IS_VARLEN=is_varlen, + num_warps=2, + num_stages=3, + ) # ============================================================================= @@ -1712,66 +1716,67 @@ def compute_ddt_dtrap_dinput_states( else: grid = (nheads, batch) - mamba3_siso_bwd_kernel_ddt_dtrap_dinput_states[grid]( - # Inputs - dscale, dgamma, dt, trap, - d_issm_state if has_input_state else dscale, # Dummy pointer if not used - input_k_state if has_input_state else dscale, - input_v_state if has_input_state else dscale, - Cu_Seqlens, - # Outputs - dDT, dTrap, - d_Input_SSM_State if has_input_state else dDT, # Dummy pointer if not used - d_Input_K_State if has_input_state else dDT, - d_Input_V_State if has_input_state else dDT, - # Strides for dScale - dscale.stride(0), dscale.stride(1), dscale.stride(2), - # Strides for dSGamma - dgamma.stride(0), dgamma.stride(1), dgamma.stride(2), - # Strides for DT - dt.stride(0), dt.stride(1), dt.stride(2), - # Strides for Trap - trap.stride(0), trap.stride(1), trap.stride(2), - # Strides for d_ISSM_State - d_issm_state.stride(0) if has_input_state else 0, - d_issm_state.stride(1) if has_input_state else 0, - d_issm_state.stride(2) if has_input_state else 0, - d_issm_state.stride(3) if has_input_state else 0, - # Strides for Input_K_State - input_k_state.stride(0) if has_input_state else 0, - input_k_state.stride(1) if has_input_state else 0, - input_k_state.stride(2) if has_input_state else 0, - # Strides for Input_V_State - input_v_state.stride(0) if has_input_state else 0, - input_v_state.stride(1) if has_input_state else 0, - input_v_state.stride(2) if has_input_state else 0, - # Stride for Cu_Seqlens - Cu_Seqlens.stride(0) if Cu_Seqlens is not None else 0, - # Strides for dDT - dDT.stride(0), dDT.stride(1), dDT.stride(2), - # Strides for dTrap - dTrap.stride(0), dTrap.stride(1), dTrap.stride(2), - # Strides for d_Input_SSM_State - d_Input_SSM_State.stride(0) if has_input_state else 0, - d_Input_SSM_State.stride(1) if has_input_state else 0, - d_Input_SSM_State.stride(2) if has_input_state else 0, - d_Input_SSM_State.stride(3) if has_input_state else 0, - # Strides for d_Input_K_State - d_Input_K_State.stride(0) if has_input_state else 0, - d_Input_K_State.stride(1) if has_input_state else 0, - d_Input_K_State.stride(2) if has_input_state else 0, - # Strides for d_Input_V_State - d_Input_V_State.stride(0) if has_input_state else 0, - d_Input_V_State.stride(1) if has_input_state else 0, - d_Input_V_State.stride(2) if has_input_state else 0, - # Dimensions - seqlen, headdim_v, headdim_qk, - # Constants - HEADDIM_V=HEADDIM_V, - HEADDIM_QK=HEADDIM_QK, - HAS_INPUT_STATE=has_input_state, - IS_VARLEN=is_varlen, - ) + with torch.cuda.device(dscale.device.index): + mamba3_siso_bwd_kernel_ddt_dtrap_dinput_states[grid]( + # Inputs + dscale, dgamma, dt, trap, + d_issm_state if has_input_state else dscale, # Dummy pointer if not used + input_k_state if has_input_state else dscale, + input_v_state if has_input_state else dscale, + Cu_Seqlens, + # Outputs + dDT, dTrap, + d_Input_SSM_State if has_input_state else dDT, # Dummy pointer if not used + d_Input_K_State if has_input_state else dDT, + d_Input_V_State if has_input_state else dDT, + # Strides for dScale + dscale.stride(0), dscale.stride(1), dscale.stride(2), + # Strides for dSGamma + dgamma.stride(0), dgamma.stride(1), dgamma.stride(2), + # Strides for DT + dt.stride(0), dt.stride(1), dt.stride(2), + # Strides for Trap + trap.stride(0), trap.stride(1), trap.stride(2), + # Strides for d_ISSM_State + d_issm_state.stride(0) if has_input_state else 0, + d_issm_state.stride(1) if has_input_state else 0, + d_issm_state.stride(2) if has_input_state else 0, + d_issm_state.stride(3) if has_input_state else 0, + # Strides for Input_K_State + input_k_state.stride(0) if has_input_state else 0, + input_k_state.stride(1) if has_input_state else 0, + input_k_state.stride(2) if has_input_state else 0, + # Strides for Input_V_State + input_v_state.stride(0) if has_input_state else 0, + input_v_state.stride(1) if has_input_state else 0, + input_v_state.stride(2) if has_input_state else 0, + # Stride for Cu_Seqlens + Cu_Seqlens.stride(0) if Cu_Seqlens is not None else 0, + # Strides for dDT + dDT.stride(0), dDT.stride(1), dDT.stride(2), + # Strides for dTrap + dTrap.stride(0), dTrap.stride(1), dTrap.stride(2), + # Strides for d_Input_SSM_State + d_Input_SSM_State.stride(0) if has_input_state else 0, + d_Input_SSM_State.stride(1) if has_input_state else 0, + d_Input_SSM_State.stride(2) if has_input_state else 0, + d_Input_SSM_State.stride(3) if has_input_state else 0, + # Strides for d_Input_K_State + d_Input_K_State.stride(0) if has_input_state else 0, + d_Input_K_State.stride(1) if has_input_state else 0, + d_Input_K_State.stride(2) if has_input_state else 0, + # Strides for d_Input_V_State + d_Input_V_State.stride(0) if has_input_state else 0, + d_Input_V_State.stride(1) if has_input_state else 0, + d_Input_V_State.stride(2) if has_input_state else 0, + # Dimensions + seqlen, headdim_v, headdim_qk, + # Constants + HEADDIM_V=HEADDIM_V, + HEADDIM_QK=HEADDIM_QK, + HAS_INPUT_STATE=has_input_state, + IS_VARLEN=is_varlen, + ) return dDT, dTrap, d_Input_SSM_State, d_Input_K_State, d_Input_V_State diff --git a/mamba_ssm/ops/triton/mamba3/mamba3_siso_fwd.py b/mamba_ssm/ops/triton/mamba3/mamba3_siso_fwd.py index 041db8eb4..3f0f7bede 100644 --- a/mamba_ssm/ops/triton/mamba3/mamba3_siso_fwd.py +++ b/mamba_ssm/ops/triton/mamba3/mamba3_siso_fwd.py @@ -629,82 +629,83 @@ def mamba3_siso_fwd( else: grid = (nheads, batch) - mamba3_siso_fwd_kernel[grid]( - # Inputs - Q, K, V, ADT, DT, Trap, Q_bias, K_bias, Angles, D, Z, - Init_SSM_State, Init_K_State, Init_V_State, cu_seqlens, - # Outputs - Out, Out_v, SSM_States, DA_CS_Store, DA_CS_SUM_Store, - Q_store, K_store, QK_store, Scale_store, Gamma_store, - Final_SSM_State, Final_K_State, - # Input strides - Q.stride(0), Q.stride(1), Q.stride(2), Q.stride(3), - K.stride(0), K.stride(1), K.stride(2), K.stride(3), - V.stride(0), V.stride(1), V.stride(2), V.stride(3), - ADT.stride(0), ADT.stride(1), ADT.stride(2), - DT.stride(0), DT.stride(1), DT.stride(2), - Trap.stride(0), Trap.stride(1), Trap.stride(2), - Q_bias.stride(0), Q_bias.stride(1), - K_bias.stride(0), K_bias.stride(1), - Angles.stride(0), Angles.stride(1), Angles.stride(2), Angles.stride(3), - D.stride(0) if D is not None else 0, - Z.stride(0) if Z is not None else 0, - Z.stride(1) if Z is not None else 0, - Z.stride(2) if Z is not None else 0, - Z.stride(3) if Z is not None else 0, - Init_SSM_State.stride(0) if Init_SSM_State is not None else 0, - Init_SSM_State.stride(1) if Init_SSM_State is not None else 0, - Init_SSM_State.stride(2) if Init_SSM_State is not None else 0, - Init_SSM_State.stride(3) if Init_SSM_State is not None else 0, - Init_K_State.stride(0) if Init_K_State is not None else 0, - Init_K_State.stride(1) if Init_K_State is not None else 0, - Init_K_State.stride(2) if Init_K_State is not None else 0, - Init_V_State.stride(0) if Init_V_State is not None else 0, - Init_V_State.stride(1) if Init_V_State is not None else 0, - Init_V_State.stride(2) if Init_V_State is not None else 0, - cu_seqlens.stride(0) if cu_seqlens is not None else 0, - # Output strides - Out.stride(0), Out.stride(1), Out.stride(2), Out.stride(3), - Out_v.stride(0) if Out_v is not None else 0, - Out_v.stride(1) if Out_v is not None else 0, - Out_v.stride(2) if Out_v is not None else 0, - Out_v.stride(3) if Out_v is not None else 0, - SSM_States.stride(0) if SSM_States is not None else 0, - SSM_States.stride(1) if SSM_States is not None else 0, - SSM_States.stride(2) if SSM_States is not None else 0, - SSM_States.stride(3) if SSM_States is not None else 0, - DA_CS_Store.stride(0) if DA_CS_Store is not None else 0, - DA_CS_Store.stride(1) if DA_CS_Store is not None else 0, - DA_CS_Store.stride(2) if DA_CS_Store is not None else 0, - DA_CS_SUM_Store.stride(0) if DA_CS_SUM_Store is not None else 0, - DA_CS_SUM_Store.stride(1) if DA_CS_SUM_Store is not None else 0, - DA_CS_SUM_Store.stride(2) if DA_CS_SUM_Store is not None else 0, - Q_store.stride(0), Q_store.stride(1), Q_store.stride(2), Q_store.stride(3), - K_store.stride(0), K_store.stride(1), K_store.stride(2), K_store.stride(3), - QK_store.stride(0), QK_store.stride(1), QK_store.stride(2), - Scale_store.stride(0), Scale_store.stride(1), Scale_store.stride(2), - Gamma_store.stride(0), Gamma_store.stride(1), Gamma_store.stride(2), - Final_SSM_State.stride(0) if Final_SSM_State is not None else 0, - Final_SSM_State.stride(1) if Final_SSM_State is not None else 0, - Final_SSM_State.stride(2) if Final_SSM_State is not None else 0, - Final_SSM_State.stride(3) if Final_SSM_State is not None else 0, - Final_K_State.stride(0) if Final_K_State is not None else 0, - Final_K_State.stride(1) if Final_K_State is not None else 0, - Final_K_State.stride(2) if Final_K_State is not None else 0, - Final_K_State.stride(3) if Final_K_State is not None else 0, - # Dimensions - seqlen, nheads_qk, headdim_qk, headdim_v, headdim_angles, - # Compile-time constants - chunk_size, - HEADDIM_QK, - HEADDIM_V, - STORE_SSM_STATES_ADT_OUTV=store_states_adt_outv, - HAS_INITIAL_STATES=Initial_States is not None, - RETURN_FINAL_STATES=return_final_states, - HAS_D=D is not None, - HAS_Z=Z is not None, - IS_VARLEN=is_varlen, - ) + with torch.cuda.device(Q.device.index): + mamba3_siso_fwd_kernel[grid]( + # Inputs + Q, K, V, ADT, DT, Trap, Q_bias, K_bias, Angles, D, Z, + Init_SSM_State, Init_K_State, Init_V_State, cu_seqlens, + # Outputs + Out, Out_v, SSM_States, DA_CS_Store, DA_CS_SUM_Store, + Q_store, K_store, QK_store, Scale_store, Gamma_store, + Final_SSM_State, Final_K_State, + # Input strides + Q.stride(0), Q.stride(1), Q.stride(2), Q.stride(3), + K.stride(0), K.stride(1), K.stride(2), K.stride(3), + V.stride(0), V.stride(1), V.stride(2), V.stride(3), + ADT.stride(0), ADT.stride(1), ADT.stride(2), + DT.stride(0), DT.stride(1), DT.stride(2), + Trap.stride(0), Trap.stride(1), Trap.stride(2), + Q_bias.stride(0), Q_bias.stride(1), + K_bias.stride(0), K_bias.stride(1), + Angles.stride(0), Angles.stride(1), Angles.stride(2), Angles.stride(3), + D.stride(0) if D is not None else 0, + Z.stride(0) if Z is not None else 0, + Z.stride(1) if Z is not None else 0, + Z.stride(2) if Z is not None else 0, + Z.stride(3) if Z is not None else 0, + Init_SSM_State.stride(0) if Init_SSM_State is not None else 0, + Init_SSM_State.stride(1) if Init_SSM_State is not None else 0, + Init_SSM_State.stride(2) if Init_SSM_State is not None else 0, + Init_SSM_State.stride(3) if Init_SSM_State is not None else 0, + Init_K_State.stride(0) if Init_K_State is not None else 0, + Init_K_State.stride(1) if Init_K_State is not None else 0, + Init_K_State.stride(2) if Init_K_State is not None else 0, + Init_V_State.stride(0) if Init_V_State is not None else 0, + Init_V_State.stride(1) if Init_V_State is not None else 0, + Init_V_State.stride(2) if Init_V_State is not None else 0, + cu_seqlens.stride(0) if cu_seqlens is not None else 0, + # Output strides + Out.stride(0), Out.stride(1), Out.stride(2), Out.stride(3), + Out_v.stride(0) if Out_v is not None else 0, + Out_v.stride(1) if Out_v is not None else 0, + Out_v.stride(2) if Out_v is not None else 0, + Out_v.stride(3) if Out_v is not None else 0, + SSM_States.stride(0) if SSM_States is not None else 0, + SSM_States.stride(1) if SSM_States is not None else 0, + SSM_States.stride(2) if SSM_States is not None else 0, + SSM_States.stride(3) if SSM_States is not None else 0, + DA_CS_Store.stride(0) if DA_CS_Store is not None else 0, + DA_CS_Store.stride(1) if DA_CS_Store is not None else 0, + DA_CS_Store.stride(2) if DA_CS_Store is not None else 0, + DA_CS_SUM_Store.stride(0) if DA_CS_SUM_Store is not None else 0, + DA_CS_SUM_Store.stride(1) if DA_CS_SUM_Store is not None else 0, + DA_CS_SUM_Store.stride(2) if DA_CS_SUM_Store is not None else 0, + Q_store.stride(0), Q_store.stride(1), Q_store.stride(2), Q_store.stride(3), + K_store.stride(0), K_store.stride(1), K_store.stride(2), K_store.stride(3), + QK_store.stride(0), QK_store.stride(1), QK_store.stride(2), + Scale_store.stride(0), Scale_store.stride(1), Scale_store.stride(2), + Gamma_store.stride(0), Gamma_store.stride(1), Gamma_store.stride(2), + Final_SSM_State.stride(0) if Final_SSM_State is not None else 0, + Final_SSM_State.stride(1) if Final_SSM_State is not None else 0, + Final_SSM_State.stride(2) if Final_SSM_State is not None else 0, + Final_SSM_State.stride(3) if Final_SSM_State is not None else 0, + Final_K_State.stride(0) if Final_K_State is not None else 0, + Final_K_State.stride(1) if Final_K_State is not None else 0, + Final_K_State.stride(2) if Final_K_State is not None else 0, + Final_K_State.stride(3) if Final_K_State is not None else 0, + # Dimensions + seqlen, nheads_qk, headdim_qk, headdim_v, headdim_angles, + # Compile-time constants + chunk_size, + HEADDIM_QK, + HEADDIM_V, + STORE_SSM_STATES_ADT_OUTV=store_states_adt_outv, + HAS_INITIAL_STATES=Initial_States is not None, + RETURN_FINAL_STATES=return_final_states, + HAS_D=D is not None, + HAS_Z=Z is not None, + IS_VARLEN=is_varlen, + ) Final_States = None if return_final_states: diff --git a/mamba_ssm/ops/triton/mamba3/mamba3_siso_step.py b/mamba_ssm/ops/triton/mamba3/mamba3_siso_step.py index 2fbde0a2f..5f2e6bbf5 100644 --- a/mamba_ssm/ops/triton/mamba3/mamba3_siso_step.py +++ b/mamba_ssm/ops/triton/mamba3/mamba3_siso_step.py @@ -333,44 +333,45 @@ def mamba3_siso_step( Output_K_State = torch.empty((batch, nheads, headdim_qk), device=device, dtype=torch.float32) grid = (nheads, batch) - mamba3_siso_step_kernel[grid]( - # Inputs - Q, K, V, ADT, DT, Trap, Q_bias, K_bias, Angles, D, Z, Input_Angle_State, Input_SSM_State, - Input_K_State, Input_V_State, - # Outputs - Out, Output_Angle_State, Output_SSM_State, Output_K_State, - # Input strides - Q.stride(0), Q.stride(1), Q.stride(2), - K.stride(0), K.stride(1), K.stride(2), - V.stride(0), V.stride(1), V.stride(2), - ADT.stride(0), ADT.stride(1), - DT.stride(0), DT.stride(1), - Trap.stride(0), Trap.stride(1), - Q_bias.stride(0), Q_bias.stride(1), - K_bias.stride(0), K_bias.stride(1), - Angles.stride(0), Angles.stride(1), Angles.stride(2), - D.stride(0) if D is not None else 0, - Z.stride(0) if Z is not None else 0, - Z.stride(1) if Z is not None else 0, - Z.stride(2) if Z is not None else 0, - Input_Angle_State.stride(0), Input_Angle_State.stride(1), Input_Angle_State.stride(2), - Input_SSM_State.stride(0), Input_SSM_State.stride(1), Input_SSM_State.stride(2), Input_SSM_State.stride(3), - Input_K_State.stride(0), Input_K_State.stride(1), Input_K_State.stride(2), - Input_V_State.stride(0), Input_V_State.stride(1), Input_V_State.stride(2), - # Output strides - Out.stride(0), Out.stride(1), Out.stride(2), - Output_Angle_State.stride(0), Output_Angle_State.stride(1), Output_Angle_State.stride(2), - Output_SSM_State.stride(0), Output_SSM_State.stride(1), Output_SSM_State.stride(2), Output_SSM_State.stride(3), - Output_K_State.stride(0), Output_K_State.stride(1), Output_K_State.stride(2), - # Dimensions - nheads_qk, - # Compile-time constants - headdim_qk, - headdim_v, - headdim_angles, - HAS_D=D is not None, - HAS_Z=Z is not None, - ) + with torch.cuda.device(Q.device.index): + mamba3_siso_step_kernel[grid]( + # Inputs + Q, K, V, ADT, DT, Trap, Q_bias, K_bias, Angles, D, Z, Input_Angle_State, Input_SSM_State, + Input_K_State, Input_V_State, + # Outputs + Out, Output_Angle_State, Output_SSM_State, Output_K_State, + # Input strides + Q.stride(0), Q.stride(1), Q.stride(2), + K.stride(0), K.stride(1), K.stride(2), + V.stride(0), V.stride(1), V.stride(2), + ADT.stride(0), ADT.stride(1), + DT.stride(0), DT.stride(1), + Trap.stride(0), Trap.stride(1), + Q_bias.stride(0), Q_bias.stride(1), + K_bias.stride(0), K_bias.stride(1), + Angles.stride(0), Angles.stride(1), Angles.stride(2), + D.stride(0) if D is not None else 0, + Z.stride(0) if Z is not None else 0, + Z.stride(1) if Z is not None else 0, + Z.stride(2) if Z is not None else 0, + Input_Angle_State.stride(0), Input_Angle_State.stride(1), Input_Angle_State.stride(2), + Input_SSM_State.stride(0), Input_SSM_State.stride(1), Input_SSM_State.stride(2), Input_SSM_State.stride(3), + Input_K_State.stride(0), Input_K_State.stride(1), Input_K_State.stride(2), + Input_V_State.stride(0), Input_V_State.stride(1), Input_V_State.stride(2), + # Output strides + Out.stride(0), Out.stride(1), Out.stride(2), + Output_Angle_State.stride(0), Output_Angle_State.stride(1), Output_Angle_State.stride(2), + Output_SSM_State.stride(0), Output_SSM_State.stride(1), Output_SSM_State.stride(2), Output_SSM_State.stride(3), + Output_K_State.stride(0), Output_K_State.stride(1), Output_K_State.stride(2), + # Dimensions + nheads_qk, + # Compile-time constants + headdim_qk, + headdim_v, + headdim_angles, + HAS_D=D is not None, + HAS_Z=Z is not None, + ) Output_States = [Output_Angle_State, Output_SSM_State, Output_K_State, V] diff --git a/mamba_ssm/ops/triton/ssd_chunk_scan.py b/mamba_ssm/ops/triton/ssd_chunk_scan.py index 7279eda98..c7bace6e6 100644 --- a/mamba_ssm/ops/triton/ssd_chunk_scan.py +++ b/mamba_ssm/ops/triton/ssd_chunk_scan.py @@ -1283,28 +1283,29 @@ def _chunk_scan_fwd(cb, x, dt, dA_cumsum, C, states, D=None, z=None, seq_idx=Non batch * nchunks, nheads) z_strides = ((z.stride(0), z.stride(1), z.stride(2), z.stride(3)) if z is not None else (0, 0, 0, 0)) - _chunk_scan_fwd_kernel[grid]( - cb, x, z, out, out_x, dt, dA_cumsum, seq_idx, C, states, D, - chunk_size, headdim, dstate, - batch, seqlen, nheads // ngroups, - cb.stride(0), cb.stride(1), cb.stride(2), cb.stride(3), cb.stride(4), - x.stride(0), x.stride(1), x.stride(2), x.stride(3), - z_strides[0], z_strides[1], z_strides[2], z_strides[3], - out.stride(0), out.stride(1), out.stride(2), out.stride(3), - dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3), - dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3), - *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)), - C.stride(0), C.stride(1), C.stride(2), C.stride(3), - states.stride(0), states.stride(1), states.stride(2), states.stride(3), states.stride(4), - D.stride(0) if D is not None else 0, - True, - D is not None, - D.dim() == 2 if D is not None else True, - BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16), - HAS_Z=z is not None, - HAS_SEQ_IDX=seq_idx is not None, - IS_TRITON_22=TRITON_22, - ) + with torch.cuda.device(x.device.index): + _chunk_scan_fwd_kernel[grid]( + cb, x, z, out, out_x, dt, dA_cumsum, seq_idx, C, states, D, + chunk_size, headdim, dstate, + batch, seqlen, nheads // ngroups, + cb.stride(0), cb.stride(1), cb.stride(2), cb.stride(3), cb.stride(4), + x.stride(0), x.stride(1), x.stride(2), x.stride(3), + z_strides[0], z_strides[1], z_strides[2], z_strides[3], + out.stride(0), out.stride(1), out.stride(2), out.stride(3), + dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3), + dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3), + *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)), + C.stride(0), C.stride(1), C.stride(2), C.stride(3), + states.stride(0), states.stride(1), states.stride(2), states.stride(3), states.stride(4), + D.stride(0) if D is not None else 0, + True, + D is not None, + D.dim() == 2 if D is not None else True, + BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16), + HAS_Z=z is not None, + HAS_SEQ_IDX=seq_idx is not None, + IS_TRITON_22=TRITON_22, + ) return out, out_x @@ -1335,28 +1336,29 @@ def _chunk_scan_fwd_wip(cb, x, dt, dA_cumsum, C, B, states, D=None, z=None, seq_ grid = lambda META: (triton.cdiv(headdim, META['BLOCK_SIZE_N']), batch * nchunks, nheads) z_strides = ((z.stride(0), z.stride(1), z.stride(2), z.stride(3)) if z is not None else (0, 0, 0, 0)) - _chunk_scan_fwd_kernel_wip[grid]( - cb, x, z, out, out_x, dt, dA_cumsum, seq_idx, C, B, states, D, - chunk_size, headdim, dstate, - batch, seqlen, nheads // ngroups, - cb.stride(0), cb.stride(1), cb.stride(2), cb.stride(3), cb.stride(4), - x.stride(0), x.stride(1), x.stride(2), x.stride(3), - z_strides[0], z_strides[1], z_strides[2], z_strides[3], - out.stride(0), out.stride(1), out.stride(2), out.stride(3), - dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3), - dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3), - *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)), - C.stride(0), C.stride(1), C.stride(2), C.stride(3), - B.stride(0), B.stride(1), B.stride(2), B.stride(3), - states.stride(0), states.stride(1), states.stride(2), states.stride(3), states.stride(4), - D.stride(0) if D is not None else 0, - D is not None, - D.dim() == 2 if D is not None else True, - BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16), - BLOCK_SIZE_M=128, - HAS_Z=z is not None, - HAS_SEQ_IDX=seq_idx is not None, - ) + with torch.cuda.device(x.device.index): + _chunk_scan_fwd_kernel_wip[grid]( + cb, x, z, out, out_x, dt, dA_cumsum, seq_idx, C, B, states, D, + chunk_size, headdim, dstate, + batch, seqlen, nheads // ngroups, + cb.stride(0), cb.stride(1), cb.stride(2), cb.stride(3), cb.stride(4), + x.stride(0), x.stride(1), x.stride(2), x.stride(3), + z_strides[0], z_strides[1], z_strides[2], z_strides[3], + out.stride(0), out.stride(1), out.stride(2), out.stride(3), + dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3), + dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3), + *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)), + C.stride(0), C.stride(1), C.stride(2), C.stride(3), + B.stride(0), B.stride(1), B.stride(2), B.stride(3), + states.stride(0), states.stride(1), states.stride(2), states.stride(3), states.stride(4), + D.stride(0) if D is not None else 0, + D is not None, + D.dim() == 2 if D is not None else True, + BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16), + BLOCK_SIZE_M=128, + HAS_Z=z is not None, + HAS_SEQ_IDX=seq_idx is not None, + ) return out, out_x