Files
rustytorch/crates/core/rtx-backend/src/ops.rs
T
2026-03-04 00:08:42 +00:00

279 lines
9.3 KiB
Rust

//! Operation traits for backends.
//!
//! These traits define operations that can be specialized per backend.
//! The main Backend trait provides basic operations, while these traits
//! provide additional specialized operations.
use crate::Backend;
/// Tensor operations that can be specialized per backend.
///
/// Most operations are defined on the main Backend trait. This trait
/// provides additional operations that may need more customization.
pub trait TensorOps<B: Backend, const D: usize>: Clone + Send + Sync {
/// Check if the tensor is contiguous in memory.
fn is_contiguous(&self) -> bool;
/// Make the tensor contiguous (copy if necessary).
fn contiguous(self) -> Self;
/// Get the number of elements.
fn numel(&self) -> usize;
/// Get the strides.
fn strides(&self) -> [usize; D];
}
/// Module operations (convolution, pooling, etc.).
pub trait ModuleOps<B: Backend>: Sized {
/// 2D convolution.
fn conv2d(
input: &B::TensorPrimitive<4>, // [batch, in_channels, height, width]
weight: &B::TensorPrimitive<4>, // [out_channels, in_channels, kH, kW]
bias: Option<&B::TensorPrimitive<1>>,
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
groups: usize,
) -> B::TensorPrimitive<4>;
/// 2D max pooling.
fn max_pool2d(
input: &B::TensorPrimitive<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
) -> B::TensorPrimitive<4>;
/// 2D average pooling.
fn avg_pool2d(
input: &B::TensorPrimitive<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
) -> B::TensorPrimitive<4>;
/// Batch normalization.
#[allow(clippy::too_many_arguments)] // Inherent to batch_norm API
fn batch_norm(
input: &B::TensorPrimitive<4>,
running_mean: &B::TensorPrimitive<1>,
running_var: &B::TensorPrimitive<1>,
weight: Option<&B::TensorPrimitive<1>>,
bias: Option<&B::TensorPrimitive<1>>,
training: bool,
momentum: f64,
eps: f64,
) -> B::TensorPrimitive<4>;
/// Dropout (returns unchanged tensor if not training).
fn dropout(input: &B::TensorPrimitive<4>, prob: f64, training: bool) -> B::TensorPrimitive<4>;
/// Linear projection (matrix multiply + optional bias).
fn linear(
input: &B::TensorPrimitive<2>,
weight: &B::TensorPrimitive<2>,
bias: Option<&B::TensorPrimitive<1>>,
) -> B::TensorPrimitive<2>;
/// Embedding lookup.
fn embedding(
indices: &B::TensorPrimitive<2>, // [batch, seq_len] as int
weight: &B::TensorPrimitive<2>, // [vocab_size, embed_dim]
) -> B::TensorPrimitive<3>; // [batch, seq_len, embed_dim]
}
/// Activation function operations.
pub trait ActivationOps<B: Backend>: Sized {
/// ReLU activation.
fn relu<const D: usize>(tensor: &B::TensorPrimitive<D>) -> B::TensorPrimitive<D>;
/// Leaky ReLU activation.
fn leaky_relu<const D: usize>(
tensor: &B::TensorPrimitive<D>,
negative_slope: f64,
) -> B::TensorPrimitive<D>;
/// Sigmoid activation.
fn sigmoid<const D: usize>(tensor: &B::TensorPrimitive<D>) -> B::TensorPrimitive<D>;
/// Tanh activation.
fn tanh<const D: usize>(tensor: &B::TensorPrimitive<D>) -> B::TensorPrimitive<D>;
/// GELU activation (Gaussian Error Linear Unit).
fn gelu<const D: usize>(tensor: &B::TensorPrimitive<D>) -> B::TensorPrimitive<D>;
/// SiLU (Swish) activation: x * sigmoid(x).
fn silu<const D: usize>(tensor: &B::TensorPrimitive<D>) -> B::TensorPrimitive<D>;
/// Mish activation: x * tanh(softplus(x)).
fn mish<const D: usize>(tensor: &B::TensorPrimitive<D>) -> B::TensorPrimitive<D>;
/// Softplus activation: log(1 + exp(x)).
fn softplus<const D: usize>(tensor: &B::TensorPrimitive<D>, beta: f64)
-> B::TensorPrimitive<D>;
/// ELU activation (Exponential Linear Unit).
fn elu<const D: usize>(tensor: &B::TensorPrimitive<D>, alpha: f64) -> B::TensorPrimitive<D>;
}
/// Attention operations - LLM performance critical.
///
/// These operations delegate to hand-optimized kernels:
/// - CUDA: FlashAttention-3 kernels
/// - Metal: Custom MSL shaders
/// - WebGPU: WGSL compute shaders
/// - CPU: Optimized BLAS-based fallback
pub trait AttentionOps<B: Backend>: Sized {
/// Flash Attention v3 (optimized scaled dot-product attention).
///
/// # Arguments
/// - `query`: `[batch, heads, seq_len, head_dim]`
/// - `key`: `[batch, kv_heads, kv_len, head_dim]`
/// - `value`: `[batch, kv_heads, kv_len, head_dim]`
/// - `mask`: Optional attention mask
/// - `scale`: Attention scale (typically 1/sqrt(head_dim))
/// - `causal`: Whether to apply causal masking
///
/// # Performance
/// - CUDA: Uses FlashAttention-3 for 5-8x speedup over naive
/// - Metal: Uses fused MSL kernel for Apple Silicon
/// - WebGPU: Uses tiled WGSL kernel
fn flash_attention(
query: &B::TensorPrimitive<4>,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
mask: Option<&B::TensorPrimitive<4>>,
scale: B::FloatElem,
causal: bool,
) -> B::TensorPrimitive<4>;
/// Flash Attention with FP8 for reduced memory.
fn flash_attention_fp8(
query: &B::TensorPrimitive<4>,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
mask: Option<&B::TensorPrimitive<4>>,
scale: B::FloatElem,
causal: bool,
) -> B::TensorPrimitive<4> {
// Default: fallback to regular flash attention
Self::flash_attention(query, key, value, mask, scale, causal)
}
/// Grouped Query Attention (GQA) - used in LLaMA 2+, Mistral.
///
/// # Arguments
/// - `query`: `[batch, q_heads, seq_len, head_dim]`
/// - `key`: `[batch, kv_heads, kv_len, head_dim]` where kv_heads < q_heads
/// - `value`: `[batch, kv_heads, kv_len, head_dim]`
///
/// Note: kv_heads divides q_heads evenly.
fn grouped_query_attention(
query: &B::TensorPrimitive<4>,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
mask: Option<&B::TensorPrimitive<4>>,
scale: B::FloatElem,
causal: bool,
num_kv_groups: usize,
) -> B::TensorPrimitive<4>;
/// Multi-Query Attention (MQA) - single KV head.
fn multi_query_attention(
query: &B::TensorPrimitive<4>,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
mask: Option<&B::TensorPrimitive<4>>,
scale: B::FloatElem,
causal: bool,
) -> B::TensorPrimitive<4> {
Self::grouped_query_attention(query, key, value, mask, scale, causal, 1)
}
/// Ring Attention for extremely long context (16M+ tokens).
///
/// Distributes attention computation across devices using ring topology.
fn ring_attention(
query: &B::TensorPrimitive<4>,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
mask: Option<&B::TensorPrimitive<4>>,
scale: B::FloatElem,
causal: bool,
chunk_size: usize,
) -> B::TensorPrimitive<4>;
/// Sliding Window Attention (used in Mistral, LongT5).
fn sliding_window_attention(
query: &B::TensorPrimitive<4>,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
window_size: usize,
scale: B::FloatElem,
) -> B::TensorPrimitive<4>;
}
/// KV-Cache operations for LLM inference.
///
/// RustyTorch++ unique feature: entropy-guided cache eviction.
pub trait KVCacheOps<B: Backend>: Sized {
/// KV-Cache handle type.
type CacheHandle: Clone + Send + Sync;
/// Create a new KV-cache.
fn create_cache(
batch_size: usize,
num_heads: usize,
max_seq_len: usize,
head_dim: usize,
device: &B::Device,
) -> Self::CacheHandle;
/// Update cache with new key-value pairs.
fn update_cache(
cache: &mut Self::CacheHandle,
key: &B::TensorPrimitive<4>,
value: &B::TensorPrimitive<4>,
position: usize,
);
/// Get cached key-value pairs.
fn get_cache(cache: &Self::CacheHandle) -> (B::TensorPrimitive<4>, B::TensorPrimitive<4>);
/// Get current sequence length in cache.
fn cache_length(cache: &Self::CacheHandle) -> usize;
/// Apply entropy-guided eviction (RustyTorch++ unique feature).
///
/// Evicts low-entropy (uninformative) tokens while preserving
/// high-entropy (important) tokens like proper nouns, numbers.
///
/// # Arguments
/// - `cache`: The cache to evict from
/// - `entropy_threshold`: Tokens below this entropy are candidates
/// - `min_retention`: Keep at least this fraction of tokens
///
/// # Returns
/// Number of tokens evicted
fn evict_by_entropy(
cache: &mut Self::CacheHandle,
entropy_threshold: f32,
min_retention: f32,
) -> usize;
/// Clear the cache.
fn clear_cache(cache: &mut Self::CacheHandle);
/// Clone the cache (for speculative decoding).
fn clone_cache(cache: &Self::CacheHandle) -> Self::CacheHandle;
/// Quantize cache to FP8 for memory savings.
fn quantize_cache_fp8(cache: &mut Self::CacheHandle);
}
#[cfg(test)]
mod tests {
// Tests would go here but require a concrete backend implementation
}