Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,86 @@
//! PTX kernel definitions for Flash Attention RTX 5090 optimization
/// Fallback PTX for backward kernel optimized for RTX 5090 Blackwell
pub const BACKWARD_FALLBACK_PTX: &str = r#"
.version 8.5
.target sm_90
.address_size 64
.visible .entry flash_backward_kernel(
.param .u64 .ptr .global .align 8 grad_output_ptr,
.param .u64 .ptr .global .align 8 q_ptr,
.param .u64 .ptr .global .align 8 k_ptr,
.param .u64 .ptr .global .align 8 v_ptr,
.param .u64 .ptr .global .align 8 output_ptr,
.param .u64 .ptr .global .align 8 lse_ptr,
.param .u64 .ptr .global .align 8 grad_q_ptr,
.param .u64 .ptr .global .align 8 grad_k_ptr,
.param .u64 .ptr .global .align 8 grad_v_ptr,
.param .u64 batch_size,
.param .u64 num_heads,
.param .u64 seq_len,
.param .u64 head_dim,
.param .f32 softmax_scale,
.param .u8 causal,
.param .u64 block_size_q,
.param .u64 block_size_kv
) {
ret;
}
.visible .entry flash_backward_dq_kernel(
.param .u64 .ptr .global .align 8 grad_output_ptr,
.param .u64 .ptr .global .align 8 q_ptr,
.param .u64 .ptr .global .align 8 k_ptr,
.param .u64 .ptr .global .align 8 v_ptr,
.param .u64 .ptr .global .align 8 lse_ptr,
.param .u64 .ptr .global .align 8 grad_q_ptr,
.param .u64 batch_size,
.param .u64 num_heads,
.param .u64 seq_len,
.param .u64 head_dim,
.param .f32 softmax_scale,
.param .u8 causal,
.param .u64 block_size_q,
.param .u64 block_size_kv
) {
ret;
}
.visible .entry flash_backward_dkv_kernel(
.param .u64 .ptr .global .align 8 grad_output_ptr,
.param .u64 .ptr .global .align 8 q_ptr,
.param .u64 .ptr .global .align 8 k_ptr,
.param .u64 .ptr .global .align 8 v_ptr,
.param .u64 .ptr .global .align 8 lse_ptr,
.param .u64 .ptr .global .align 8 grad_k_ptr,
.param .u64 .ptr .global .align 8 grad_v_ptr,
.param .u64 batch_size,
.param .u64 num_heads,
.param .u64 seq_len,
.param .u64 head_dim,
.param .f32 softmax_scale,
.param .u8 causal,
.param .u64 block_size_q,
.param .u64 block_size_kv
) {
ret;
}
"#;
/// PTX loader with build-time and fallback support
pub struct PtxLoader;
impl PtxLoader {
/// Load backward PTX with intelligent fallback for RTX 5090
pub fn load_backward_ptx() -> crate::error::FlashResult<String> {
// For now, always use fallback PTX to ensure RTX 5090 sm_90 compatibility
tracing::warn!("Using fallback PTX for sm_90 RTX 5090 compatibility");
Ok(BACKWARD_FALLBACK_PTX.to_string())
}
/// Validate PTX for RTX 5090 compatibility
pub fn validate_ptx(ptx: &str) -> bool {
ptx.contains(".target sm_90") && ptx.contains(".version 8.5")
}
}