GPU Mamba forward + rtx-tensor device-pointer API + matrix exponential
Performance Benchmarks / Run Benchmarks (pull_request) Has been cancelled
CI / Format Check (pull_request) Has been cancelled
Documentation / Build User Guide (pull_request) Has been cancelled
CI / Clippy Check (pull_request) Has been cancelled
CI / Build (macos-latest) (pull_request) Has been cancelled
CI / Build (ubuntu-latest) (pull_request) Has been cancelled
CI / Test (macos-latest) (pull_request) Has been cancelled
CI / Test (ubuntu-latest) (pull_request) Has been cancelled
CI / Build CPU-Only (Explicit) (pull_request) Has been cancelled
CI / CI Success (pull_request) Has been cancelled
Documentation / Build API Documentation (pull_request) Has been cancelled
GPU Tests / Check GPU Availability (pull_request) Has been cancelled
GPU Tests / CUDA Tests (11.8) (pull_request) Has been cancelled
GPU Tests / CUDA Tests (12.1) (pull_request) Has been cancelled
GPU Tests / Metal Tests (pull_request) Has been cancelled
Performance Benchmarks / Run Benchmarks (pull_request) Has been cancelled
CI / Format Check (pull_request) Has been cancelled
Documentation / Build User Guide (pull_request) Has been cancelled
CI / Clippy Check (pull_request) Has been cancelled
CI / Build (macos-latest) (pull_request) Has been cancelled
CI / Build (ubuntu-latest) (pull_request) Has been cancelled
CI / Test (macos-latest) (pull_request) Has been cancelled
CI / Test (ubuntu-latest) (pull_request) Has been cancelled
CI / Build CPU-Only (Explicit) (pull_request) Has been cancelled
CI / CI Success (pull_request) Has been cancelled
Documentation / Build API Documentation (pull_request) Has been cancelled
GPU Tests / Check GPU Availability (pull_request) Has been cancelled
GPU Tests / CUDA Tests (11.8) (pull_request) Has been cancelled
GPU Tests / CUDA Tests (12.1) (pull_request) Has been cancelled
GPU Tests / Metal Tests (pull_request) Has been cancelled
Persist accumulated WIP across rtx-tensor / rtx-transformers / rtx-runtime. Two related feature groups: GPU enablement (unblocks the Phase-3 spec §8 CUDA Mamba path): - rtx-tensor: `Tensor::cuda_device_ptr()` + storage GPU-buffer accessors (`raw_ptr.rs`) expose the raw CUdeviceptr that kernel launches need — the "rtx-tensor GPU memory access API" the Mamba CUDA kernels were blocked on. - rtx-transformers: `MambaBlock::forward_cuda` runs the four linear projections through cuBLAS on-device (in/out/x/dt_proj), keeping the selective scan + conv1d + activations on CPU; dispatched automatically from `forward` when on a CUDA device under the `cuda` feature. Updated `mamba_cuda_kernels.rs` accordingly. - supporting plumbing in rtx-runtime stream/bridge and rtx-tensor storage/conversion/concatenation/creation + rtx-flash-attention. Linear algebra (rtx-tensor): - `linalg/matrix_exp.rs`: real matrix exponential via scaling-and-squaring with a degree-13 Padé approximant (Higham 2005), f64 internally. - `complex/linalg.rs`: complex matmul/adjoint, Hermitian eigendecomposition (`ComplexEigenResult`), and the complex matrix exponential, nalgebra-backed. - tests for both. Builds verified on the CPU path (`cargo check -p rtx-tensor -p rtx-transformers -p rtx-runtime -p rtx-flash-attention` clean). The `cuda` feature and the rtx-backend-cuda NVCC build remain unbuildable on this host (CUDA/glibc header mismatch) — pre-existing and unrelated to these changes.
This commit is contained in:
@@ -6,7 +6,7 @@
|
|||||||
use super::event::Event;
|
use super::event::Event;
|
||||||
use super::types::{BackendType, DeviceId, StreamId, StreamPriority, StreamStats};
|
use super::types::{BackendType, DeviceId, StreamId, StreamPriority, StreamStats};
|
||||||
use crate::allocator::DevicePtr;
|
use crate::allocator::DevicePtr;
|
||||||
use crate::error::Result;
|
use crate::error::{Result, RuntimeError};
|
||||||
use parking_lot::Mutex;
|
use parking_lot::Mutex;
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
use std::sync::atomic::{AtomicU64, Ordering};
|
use std::sync::atomic::{AtomicU64, Ordering};
|
||||||
|
|||||||
@@ -5,6 +5,7 @@
|
|||||||
|
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
use crate::cuda_backend::CudaStreamHandle;
|
use crate::cuda_backend::CudaStreamHandle;
|
||||||
|
use crate::BackendType;
|
||||||
use crate::device::{Stream, StreamId};
|
use crate::device::{Stream, StreamId};
|
||||||
use crate::error::{Result, RuntimeError};
|
use crate::error::{Result, RuntimeError};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|||||||
@@ -0,0 +1,264 @@
|
|||||||
|
//! Linear algebra for complex tensors: matrix multiplication, adjoint,
|
||||||
|
//! Hermitian eigendecomposition, and the matrix exponential.
|
||||||
|
//!
|
||||||
|
//! These are the operations quantum-style workloads (Hamiltonians, unitary
|
||||||
|
//! propagators) need on top of the elementwise complex arithmetic. The
|
||||||
|
//! decompositions run internally in f64 via nalgebra for numerical accuracy
|
||||||
|
//! and convert back to the tensor dtype at the boundary.
|
||||||
|
|
||||||
|
use nalgebra::{Complex, DMatrix};
|
||||||
|
|
||||||
|
use super::core::ComplexTensor;
|
||||||
|
use super::traits::ComplexFloat;
|
||||||
|
use crate::{Result, Tensor, TensorError};
|
||||||
|
|
||||||
|
/// Result of a Hermitian eigendecomposition A = V·Λ·V†.
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct ComplexEigenResult<T: ComplexFloat> {
|
||||||
|
/// Real eigenvalues in ascending order, shape `[n]`.
|
||||||
|
pub eigenvalues: Tensor,
|
||||||
|
/// Unit eigenvectors as matrix columns, shape `[n, n]` (if requested).
|
||||||
|
pub eigenvectors: Option<ComplexTensor<T>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T: ComplexFloat> ComplexTensor<T> {
|
||||||
|
/// Matrix multiplication of two complex matrices:
|
||||||
|
/// (A + iB)(C + iD) = (AC − BD) + i(AD + BC), using four real matmuls so
|
||||||
|
/// the computation stays on the tensor backend.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
/// Returns an error if either operand is not 2-D or the inner dimensions
|
||||||
|
/// do not match (propagated from the real matmuls).
|
||||||
|
pub fn matmul(&self, other: &Self) -> Result<Self> {
|
||||||
|
let ac = self.real().matmul(other.real())?;
|
||||||
|
let bd = self.imag().matmul(other.imag())?;
|
||||||
|
let ad = self.real().matmul(other.imag())?;
|
||||||
|
let bc = self.imag().matmul(other.real())?;
|
||||||
|
Self::from_real_imag(ac.sub(&bd)?, ad.add(&bc)?)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Conjugate transpose A† of a 2-D complex matrix.
|
||||||
|
///
|
||||||
|
/// The transpose is materialized (data copy) rather than expressed as a
|
||||||
|
/// strided view: `Tensor::transpose` returns a stride-swapped view, and
|
||||||
|
/// not all consumers (`to_vec`, `matmul`) honor non-contiguous strides.
|
||||||
|
pub fn adjoint(&self) -> Result<Self> {
|
||||||
|
let dims = self.shape().dims().to_vec();
|
||||||
|
if dims.len() != 2 {
|
||||||
|
return Err(TensorError::shape(format!(
|
||||||
|
"adjoint requires a 2-D matrix, got shape {dims:?}"
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
let (rows, cols) = (dims[0], dims[1]);
|
||||||
|
let re = self.real().to_vec()?;
|
||||||
|
let im = self.imag().to_vec()?;
|
||||||
|
let mut re_t = vec![0.0_f32; rows * cols];
|
||||||
|
let mut im_t = vec![0.0_f32; rows * cols];
|
||||||
|
for i in 0..rows {
|
||||||
|
for j in 0..cols {
|
||||||
|
re_t[j * rows + i] = re[i * cols + j];
|
||||||
|
im_t[j * rows + i] = -im[i * cols + j];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let device = self.device().clone();
|
||||||
|
let real = Tensor::from_data(re_t, vec![cols, rows], &device)?;
|
||||||
|
let imag = Tensor::from_data(im_t, vec![cols, rows], &device)?;
|
||||||
|
Self::from_real_imag(real, imag)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Eigendecomposition of a Hermitian matrix: A = V·Λ·V† with real
|
||||||
|
/// eigenvalues in ascending order.
|
||||||
|
///
|
||||||
|
/// The matrix is symmetrized as (A + A†)/2 before decomposition, so small
|
||||||
|
/// numerical deviations from Hermiticity are tolerated. Internally the
|
||||||
|
/// problem is lifted to the real symmetric 2n×2n block form
|
||||||
|
/// `[[Re, −Im], [Im, Re]]` (each eigenvalue appears exactly twice) and
|
||||||
|
/// solved in f64.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
/// Returns an error if the tensor is not a square 2-D matrix.
|
||||||
|
pub fn eigh(&self, eigenvectors: bool) -> Result<ComplexEigenResult<T>> {
|
||||||
|
let (re, im) = self.to_nalgebra()?;
|
||||||
|
let n = re.nrows();
|
||||||
|
if re.ncols() != n {
|
||||||
|
return Err(TensorError::shape(format!(
|
||||||
|
"eigh requires a square matrix, got {}×{}",
|
||||||
|
n,
|
||||||
|
re.ncols()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hermitianize: H = (A + A†)/2.
|
||||||
|
let h_re = (&re + re.transpose()) * 0.5;
|
||||||
|
let h_im = (&im - im.transpose()) * 0.5;
|
||||||
|
|
||||||
|
// Real symmetric embedding [[Re, −Im], [Im, Re]].
|
||||||
|
let mut big = DMatrix::<f64>::zeros(2 * n, 2 * n);
|
||||||
|
for i in 0..n {
|
||||||
|
for j in 0..n {
|
||||||
|
big[(i, j)] = h_re[(i, j)];
|
||||||
|
big[(i + n, j + n)] = h_re[(i, j)];
|
||||||
|
big[(i, j + n)] = -h_im[(i, j)];
|
||||||
|
big[(i + n, j)] = h_im[(i, j)];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let eig = big.symmetric_eigen();
|
||||||
|
|
||||||
|
// Sort the 2n eigenpairs ascending; every physical eigenvalue appears
|
||||||
|
// exactly twice, so taking every other sorted entry yields the n
|
||||||
|
// unique values with correct multiplicities.
|
||||||
|
let mut order: Vec<usize> = (0..2 * n).collect();
|
||||||
|
order.sort_by(|&a, &b| eig.eigenvalues[a].total_cmp(&eig.eigenvalues[b]));
|
||||||
|
|
||||||
|
let mut values = Vec::with_capacity(n);
|
||||||
|
let mut vec_re = DMatrix::<f64>::zeros(n, n);
|
||||||
|
let mut vec_im = DMatrix::<f64>::zeros(n, n);
|
||||||
|
for k in 0..n {
|
||||||
|
let idx = order[2 * k];
|
||||||
|
values.push(eig.eigenvalues[idx] as f32);
|
||||||
|
if eigenvectors {
|
||||||
|
let col = eig.eigenvectors.column(idx);
|
||||||
|
// [x; y] ↔ v = x + iy (unit 2n-vector ⇒ unit complex vector).
|
||||||
|
for i in 0..n {
|
||||||
|
vec_re[(i, k)] = col[i];
|
||||||
|
vec_im[(i, k)] = col[i + n];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let device = self.device().clone();
|
||||||
|
let eigenvalues = Tensor::from_data(values, vec![n], &device)?;
|
||||||
|
let eigenvectors = if eigenvectors {
|
||||||
|
Some(Self::from_nalgebra(&vec_re, &vec_im, &device)?)
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
Ok(ComplexEigenResult {
|
||||||
|
eigenvalues,
|
||||||
|
eigenvectors,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Matrix exponential exp(A) of a square complex matrix via
|
||||||
|
/// scaling-and-squaring with a degree-13 Padé approximant (Higham 2005),
|
||||||
|
/// computed internally in f64.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
/// Returns an error if the tensor is not a square 2-D matrix.
|
||||||
|
pub fn matrix_exp(&self) -> Result<Self> {
|
||||||
|
let (re, im) = self.to_nalgebra()?;
|
||||||
|
let n = re.nrows();
|
||||||
|
if re.ncols() != n {
|
||||||
|
return Err(TensorError::shape(format!(
|
||||||
|
"matrix_exp requires a square matrix, got {}×{}",
|
||||||
|
n,
|
||||||
|
re.ncols()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
let a = DMatrix::<Complex<f64>>::from_fn(n, n, |i, j| Complex::new(re[(i, j)], im[(i, j)]));
|
||||||
|
let e = expm_pade13(&a)?;
|
||||||
|
let e_re = DMatrix::<f64>::from_fn(n, n, |i, j| e[(i, j)].re);
|
||||||
|
let e_im = DMatrix::<f64>::from_fn(n, n, |i, j| e[(i, j)].im);
|
||||||
|
Self::from_nalgebra(&e_re, &e_im, &self.device().clone())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Read the matrix into f64 nalgebra storage (row-major tensor layout).
|
||||||
|
fn to_nalgebra(&self) -> Result<(DMatrix<f64>, DMatrix<f64>)> {
|
||||||
|
let dims = self.shape().dims();
|
||||||
|
if dims.len() != 2 {
|
||||||
|
return Err(TensorError::shape(format!(
|
||||||
|
"expected a 2-D matrix, got shape {dims:?}"
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
let (rows, cols) = (dims[0], dims[1]);
|
||||||
|
let re_data: Vec<f64> = self.real().to_vec()?.iter().map(|&x| x as f64).collect();
|
||||||
|
let im_data: Vec<f64> = self.imag().to_vec()?.iter().map(|&x| x as f64).collect();
|
||||||
|
if re_data.len() != rows * cols || im_data.len() != rows * cols {
|
||||||
|
return Err(TensorError::shape(
|
||||||
|
"tensor data length does not match its shape".to_string(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
Ok((
|
||||||
|
DMatrix::from_row_slice(rows, cols, &re_data),
|
||||||
|
DMatrix::from_row_slice(rows, cols, &im_data),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build a complex tensor from f64 nalgebra parts.
|
||||||
|
fn from_nalgebra(
|
||||||
|
re: &DMatrix<f64>,
|
||||||
|
im: &DMatrix<f64>,
|
||||||
|
device: &crate::Device,
|
||||||
|
) -> Result<Self> {
|
||||||
|
let (rows, cols) = (re.nrows(), re.ncols());
|
||||||
|
let mut re_data = Vec::with_capacity(rows * cols);
|
||||||
|
let mut im_data = Vec::with_capacity(rows * cols);
|
||||||
|
for i in 0..rows {
|
||||||
|
for j in 0..cols {
|
||||||
|
re_data.push(re[(i, j)] as f32);
|
||||||
|
im_data.push(im[(i, j)] as f32);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let real = Tensor::from_data(re_data, vec![rows, cols], device)?;
|
||||||
|
let imag = Tensor::from_data(im_data, vec![rows, cols], device)?;
|
||||||
|
Self::from_real_imag(real, imag)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Scaling-and-squaring Padé-13 matrix exponential on complex f64 matrices.
|
||||||
|
fn expm_pade13(a: &DMatrix<Complex<f64>>) -> Result<DMatrix<Complex<f64>>> {
|
||||||
|
const THETA_13: f64 = 5.371_920_351_148_152;
|
||||||
|
const B: [f64; 14] = [
|
||||||
|
64_764_752_532_480_000.0,
|
||||||
|
32_382_376_266_240_000.0,
|
||||||
|
7_771_770_303_897_600.0,
|
||||||
|
1_187_353_796_428_800.0,
|
||||||
|
129_060_195_264_000.0,
|
||||||
|
10_559_470_521_600.0,
|
||||||
|
670_442_572_800.0,
|
||||||
|
33_522_128_640.0,
|
||||||
|
1_323_241_920.0,
|
||||||
|
40_840_800.0,
|
||||||
|
960_960.0,
|
||||||
|
16_380.0,
|
||||||
|
182.0,
|
||||||
|
1.0,
|
||||||
|
];
|
||||||
|
|
||||||
|
let n = a.nrows();
|
||||||
|
// 1-norm (maximum absolute column sum).
|
||||||
|
let norm = (0..n)
|
||||||
|
.map(|j| (0..n).map(|i| a[(i, j)].norm_sqr().sqrt()).sum::<f64>())
|
||||||
|
.fold(0.0_f64, f64::max);
|
||||||
|
let s = if norm > THETA_13 {
|
||||||
|
(norm / THETA_13).log2().ceil() as u32
|
||||||
|
} else {
|
||||||
|
0
|
||||||
|
};
|
||||||
|
let scale = Complex::new((0.5_f64).powi(s as i32), 0.0);
|
||||||
|
let a1 = a * scale;
|
||||||
|
|
||||||
|
let id = DMatrix::<Complex<f64>>::identity(n, n);
|
||||||
|
let a2 = &a1 * &a1;
|
||||||
|
let a4 = &a2 * &a2;
|
||||||
|
let a6 = &a2 * &a4;
|
||||||
|
|
||||||
|
let c = |k: usize| Complex::new(B[k], 0.0);
|
||||||
|
let u_inner = &a6 * c(13) + &a4 * c(11) + &a2 * c(9);
|
||||||
|
let u_poly = &a6 * &u_inner + &a6 * c(7) + &a4 * c(5) + &a2 * c(3) + &id * c(1);
|
||||||
|
let u = &a1 * &u_poly;
|
||||||
|
let v_inner = &a6 * c(12) + &a4 * c(10) + &a2 * c(8);
|
||||||
|
let v = &a6 * &v_inner + &a6 * c(6) + &a4 * c(4) + &a2 * c(2) + &id * c(0);
|
||||||
|
|
||||||
|
// exp(A) ≈ (V − U)⁻¹(V + U).
|
||||||
|
let p = &v + &u;
|
||||||
|
let q = &v - &u;
|
||||||
|
let mut e = q
|
||||||
|
.lu()
|
||||||
|
.solve(&p)
|
||||||
|
.ok_or_else(|| TensorError::numerical("Padé denominator is singular".to_string()))?;
|
||||||
|
for _ in 0..s {
|
||||||
|
e = &e * &e;
|
||||||
|
}
|
||||||
|
Ok(e)
|
||||||
|
}
|
||||||
@@ -0,0 +1,204 @@
|
|||||||
|
//! Tests for complex matrix linear algebra: matmul, adjoint, eigh, matrix_exp.
|
||||||
|
|
||||||
|
use crate::complex::ComplexTensor;
|
||||||
|
use crate::{Device, Tensor};
|
||||||
|
|
||||||
|
fn complex_from_rows(
|
||||||
|
re: Vec<f32>,
|
||||||
|
im: Vec<f32>,
|
||||||
|
n: usize,
|
||||||
|
device: &Device,
|
||||||
|
) -> ComplexTensor<f32> {
|
||||||
|
let real = Tensor::from_data(re, vec![n, n], device).unwrap();
|
||||||
|
let imag = Tensor::from_data(im, vec![n, n], device).unwrap();
|
||||||
|
ComplexTensor::from_real_imag(real, imag).unwrap()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn assert_close(a: f32, b: f32, tol: f32, what: &str) {
|
||||||
|
assert!((a - b).abs() < tol, "{what}: {a} vs {b}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn matmul_matches_hand_computed_product() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// A = [[1+i, 2], [0, 1-i]], B = [[i, 0], [1, 1]]
|
||||||
|
let a = complex_from_rows(
|
||||||
|
vec![1.0, 2.0, 0.0, 1.0],
|
||||||
|
vec![1.0, 0.0, 0.0, -1.0],
|
||||||
|
2,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let b = complex_from_rows(
|
||||||
|
vec![0.0, 0.0, 1.0, 1.0],
|
||||||
|
vec![1.0, 0.0, 0.0, 0.0],
|
||||||
|
2,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let c = a.matmul(&b).unwrap();
|
||||||
|
let re = c.real().to_vec().unwrap();
|
||||||
|
let im = c.imag().to_vec().unwrap();
|
||||||
|
// C[0][0] = (1+i)·i + 2·1 = i + i² + 2 = 1 + i
|
||||||
|
assert_close(re[0], 1.0, 1e-6, "C00 re");
|
||||||
|
assert_close(im[0], 1.0, 1e-6, "C00 im");
|
||||||
|
// C[0][1] = (1+i)·0 + 2·1 = 2
|
||||||
|
assert_close(re[1], 2.0, 1e-6, "C01 re");
|
||||||
|
assert_close(im[1], 0.0, 1e-6, "C01 im");
|
||||||
|
// C[1][0] = 0·i + (1−i)·1 = 1 − i
|
||||||
|
assert_close(re[2], 1.0, 1e-6, "C10 re");
|
||||||
|
assert_close(im[2], -1.0, 1e-6, "C10 im");
|
||||||
|
// C[1][1] = (1−i)
|
||||||
|
assert_close(re[3], 1.0, 1e-6, "C11 re");
|
||||||
|
assert_close(im[3], -1.0, 1e-6, "C11 im");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn adjoint_of_product_reverses_order() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
let a = complex_from_rows(
|
||||||
|
vec![1.0, 2.0, -1.0, 0.5],
|
||||||
|
vec![0.5, -1.0, 2.0, 0.0],
|
||||||
|
2,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let b = complex_from_rows(
|
||||||
|
vec![0.0, 1.0, 1.0, -2.0],
|
||||||
|
vec![1.0, 0.0, -0.5, 1.0],
|
||||||
|
2,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let lhs = a.matmul(&b).unwrap().adjoint().unwrap();
|
||||||
|
let rhs = b.adjoint().unwrap().matmul(&a.adjoint().unwrap()).unwrap();
|
||||||
|
let (lr, li) = (lhs.real().to_vec().unwrap(), lhs.imag().to_vec().unwrap());
|
||||||
|
let (rr, ri) = (rhs.real().to_vec().unwrap(), rhs.imag().to_vec().unwrap());
|
||||||
|
for k in 0..4 {
|
||||||
|
assert_close(lr[k], rr[k], 1e-5, "(AB)† vs B†A† re");
|
||||||
|
assert_close(li[k], ri[k], 1e-5, "(AB)† vs B†A† im");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn eigh_pauli_y_eigenvalues_are_plus_minus_one() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// Y = [[0, −i], [i, 0]]: genuinely complex Hermitian, eigenvalues ±1.
|
||||||
|
let y = complex_from_rows(
|
||||||
|
vec![0.0, 0.0, 0.0, 0.0],
|
||||||
|
vec![0.0, -1.0, 1.0, 0.0],
|
||||||
|
2,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let result = y.eigh(true).unwrap();
|
||||||
|
let vals = result.eigenvalues.to_vec().unwrap();
|
||||||
|
assert_close(vals[0], -1.0, 1e-5, "λ0");
|
||||||
|
assert_close(vals[1], 1.0, 1e-5, "λ1");
|
||||||
|
|
||||||
|
// Residual check: Y·v = λ·v for each column.
|
||||||
|
let v = result.eigenvectors.unwrap();
|
||||||
|
let yv = y.matmul(&v).unwrap();
|
||||||
|
let (yv_re, yv_im) = (yv.real().to_vec().unwrap(), yv.imag().to_vec().unwrap());
|
||||||
|
let (v_re, v_im) = (v.real().to_vec().unwrap(), v.imag().to_vec().unwrap());
|
||||||
|
for col in 0..2 {
|
||||||
|
let lam = vals[col];
|
||||||
|
for row in 0..2 {
|
||||||
|
let idx = row * 2 + col;
|
||||||
|
assert_close(yv_re[idx], lam * v_re[idx], 1e-5, "Y·v re");
|
||||||
|
assert_close(yv_im[idx], lam * v_im[idx], 1e-5, "Y·v im");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn eigh_eigenvectors_are_orthonormal() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// A fixed 3×3 Hermitian matrix with complex off-diagonals.
|
||||||
|
let h = complex_from_rows(
|
||||||
|
vec![2.0, 0.5, 0.0, 0.5, -1.0, 1.0, 0.0, 1.0, 0.5],
|
||||||
|
vec![0.0, 0.3, -0.2, -0.3, 0.0, 0.4, 0.2, -0.4, 0.0],
|
||||||
|
3,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let result = h.eigh(true).unwrap();
|
||||||
|
let v = result.eigenvectors.unwrap();
|
||||||
|
let gram = v.adjoint().unwrap().matmul(&v).unwrap();
|
||||||
|
let (g_re, g_im) = (gram.real().to_vec().unwrap(), gram.imag().to_vec().unwrap());
|
||||||
|
for i in 0..3 {
|
||||||
|
for j in 0..3 {
|
||||||
|
let expected = if i == j { 1.0 } else { 0.0 };
|
||||||
|
assert_close(g_re[i * 3 + j], expected, 1e-5, "V†V re");
|
||||||
|
assert_close(g_im[i * 3 + j], 0.0, 1e-5, "V†V im");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Eigenvalues ascending.
|
||||||
|
let vals = result.eigenvalues.to_vec().unwrap();
|
||||||
|
assert!(vals[0] <= vals[1] && vals[1] <= vals[2]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn matrix_exp_of_zero_is_identity() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
let z = complex_from_rows(vec![0.0; 4], vec![0.0; 4], 2, &device);
|
||||||
|
let e = z.matrix_exp().unwrap();
|
||||||
|
let (re, im) = (e.real().to_vec().unwrap(), e.imag().to_vec().unwrap());
|
||||||
|
for i in 0..2 {
|
||||||
|
for j in 0..2 {
|
||||||
|
let expected = if i == j { 1.0 } else { 0.0 };
|
||||||
|
assert_close(re[i * 2 + j], expected, 1e-6, "exp(0) re");
|
||||||
|
assert_close(im[i * 2 + j], 0.0, 1e-6, "exp(0) im");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn matrix_exp_of_skew_hermitian_is_unitary() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// exp(−iθY): U = [[cosθ, −sinθ], [sinθ, cosθ]] for −iθY real form…
|
||||||
|
// verify the general property U†U = I instead of a specific matrix.
|
||||||
|
let theta = 0.7_f32;
|
||||||
|
// −iθY = [[0, −θ], [θ, 0]] (purely real skew-symmetric here).
|
||||||
|
let a = complex_from_rows(
|
||||||
|
vec![0.0, -theta, theta, 0.0],
|
||||||
|
vec![0.0; 4],
|
||||||
|
2,
|
||||||
|
&device,
|
||||||
|
);
|
||||||
|
let u = a.matrix_exp().unwrap();
|
||||||
|
let gram = u.adjoint().unwrap().matmul(&u).unwrap();
|
||||||
|
let (g_re, g_im) = (gram.real().to_vec().unwrap(), gram.imag().to_vec().unwrap());
|
||||||
|
for i in 0..2 {
|
||||||
|
for j in 0..2 {
|
||||||
|
let expected = if i == j { 1.0 } else { 0.0 };
|
||||||
|
assert_close(g_re[i * 2 + j], expected, 1e-5, "U†U re");
|
||||||
|
assert_close(g_im[i * 2 + j], 0.0, 1e-5, "U†U im");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// And the rotation form: U00 = cosθ, U10 = sinθ.
|
||||||
|
let u_re = u.real().to_vec().unwrap();
|
||||||
|
assert_close(u_re[0], theta.cos(), 1e-5, "cosθ");
|
||||||
|
assert_close(u_re[2], theta.sin(), 1e-5, "sinθ");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn matrix_exp_diagonal_imaginary_gives_phases() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// exp(i·diag(φ1, φ2)) = diag(e^{iφ1}, e^{iφ2}).
|
||||||
|
let (p1, p2) = (0.3_f32, -1.1_f32);
|
||||||
|
let a = complex_from_rows(vec![0.0; 4], vec![p1, 0.0, 0.0, p2], 2, &device);
|
||||||
|
let e = a.matrix_exp().unwrap();
|
||||||
|
let (re, im) = (e.real().to_vec().unwrap(), e.imag().to_vec().unwrap());
|
||||||
|
assert_close(re[0], p1.cos(), 1e-6, "e^{iφ1} re");
|
||||||
|
assert_close(im[0], p1.sin(), 1e-6, "e^{iφ1} im");
|
||||||
|
assert_close(re[3], p2.cos(), 1e-6, "e^{iφ2} re");
|
||||||
|
assert_close(im[3], p2.sin(), 1e-6, "e^{iφ2} im");
|
||||||
|
assert_close(re[1], 0.0, 1e-6, "off-diag");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn matrix_exp_large_norm_uses_scaling_correctly() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// Norm ≫ θ13 forces the scaling-and-squaring path:
|
||||||
|
// exp(diag(10, −10)) = diag(e^10, e^−10).
|
||||||
|
let a = complex_from_rows(vec![10.0, 0.0, 0.0, -10.0], vec![0.0; 4], 2, &device);
|
||||||
|
let e = a.matrix_exp().unwrap();
|
||||||
|
let re = e.real().to_vec().unwrap();
|
||||||
|
assert!((re[0] - 10.0_f32.exp()).abs() / 10.0_f32.exp() < 1e-4);
|
||||||
|
assert!((re[3] - (-10.0_f32).exp()).abs() < 1e-6);
|
||||||
|
}
|
||||||
@@ -47,6 +47,7 @@ pub mod autograd;
|
|||||||
pub mod complex_ops;
|
pub mod complex_ops;
|
||||||
pub mod core;
|
pub mod core;
|
||||||
pub mod fft;
|
pub mod fft;
|
||||||
|
pub mod linalg;
|
||||||
pub mod tensor;
|
pub mod tensor;
|
||||||
pub mod traits;
|
pub mod traits;
|
||||||
|
|
||||||
@@ -71,9 +72,13 @@ mod simple_conjugate_tests;
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod phase_magnitude_test;
|
mod phase_magnitude_test;
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod linalg_tests;
|
||||||
|
|
||||||
pub mod phase_magnitude_demo;
|
pub mod phase_magnitude_demo;
|
||||||
|
|
||||||
// Re-export main types
|
// Re-export main types
|
||||||
|
pub use linalg::ComplexEigenResult;
|
||||||
pub use tensor::{ComplexFloat, ComplexTensor, NormMode};
|
pub use tensor::{ComplexFloat, ComplexTensor, NormMode};
|
||||||
|
|
||||||
/// Type alias for f32 complex tensors
|
/// Type alias for f32 complex tensors
|
||||||
|
|||||||
@@ -0,0 +1,107 @@
|
|||||||
|
//! Matrix exponential for real square matrices.
|
||||||
|
//!
|
||||||
|
//! Scaling-and-squaring with a degree-13 Padé approximant (Higham 2005),
|
||||||
|
//! computed internally in f64 via nalgebra for accuracy and converted back
|
||||||
|
//! to the tensor dtype at the boundary. The complex-matrix variant lives on
|
||||||
|
//! `ComplexTensor::matrix_exp`.
|
||||||
|
|
||||||
|
use nalgebra::DMatrix;
|
||||||
|
|
||||||
|
use crate::{Result, Tensor, TensorError};
|
||||||
|
|
||||||
|
impl Tensor {
|
||||||
|
/// Matrix exponential exp(A) of a square 2-D matrix.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
/// Returns an error if the tensor is not a square 2-D matrix or the Padé
|
||||||
|
/// denominator turns out singular (pathological input).
|
||||||
|
///
|
||||||
|
/// # Example
|
||||||
|
/// ```
|
||||||
|
/// use rtx_tensor::{Tensor, Device};
|
||||||
|
/// let device = Device::cpu();
|
||||||
|
/// // Nilpotent: exp([[0,1],[0,0]]) = [[1,1],[0,1]].
|
||||||
|
/// let a = Tensor::from_data(vec![0.0, 1.0, 0.0, 0.0], vec![2, 2], &device).unwrap();
|
||||||
|
/// let e = a.matrix_exp().unwrap();
|
||||||
|
/// let v = e.to_vec().unwrap();
|
||||||
|
/// assert!((v[0] - 1.0).abs() < 1e-6 && (v[1] - 1.0).abs() < 1e-6);
|
||||||
|
/// ```
|
||||||
|
pub fn matrix_exp(&self) -> Result<Tensor> {
|
||||||
|
let dims = self.shape().dims().to_vec();
|
||||||
|
if dims.len() != 2 || dims[0] != dims[1] {
|
||||||
|
return Err(TensorError::shape(format!(
|
||||||
|
"matrix_exp requires a square 2-D matrix, got shape {dims:?}"
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
let n = dims[0];
|
||||||
|
let data: Vec<f64> = self.to_vec()?.iter().map(|&x| x as f64).collect();
|
||||||
|
if data.len() != n * n {
|
||||||
|
return Err(TensorError::shape(
|
||||||
|
"tensor data length does not match its shape".to_string(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let a = DMatrix::from_row_slice(n, n, &data);
|
||||||
|
let e = expm_pade13_real(&a)?;
|
||||||
|
let mut out = Vec::with_capacity(n * n);
|
||||||
|
for i in 0..n {
|
||||||
|
for j in 0..n {
|
||||||
|
out.push(e[(i, j)] as f32);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Tensor::from_data(out, vec![n, n], self.device())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Scaling-and-squaring Padé-13 on real f64 matrices.
|
||||||
|
fn expm_pade13_real(a: &DMatrix<f64>) -> Result<DMatrix<f64>> {
|
||||||
|
const THETA_13: f64 = 5.371_920_351_148_152;
|
||||||
|
const B: [f64; 14] = [
|
||||||
|
64_764_752_532_480_000.0,
|
||||||
|
32_382_376_266_240_000.0,
|
||||||
|
7_771_770_303_897_600.0,
|
||||||
|
1_187_353_796_428_800.0,
|
||||||
|
129_060_195_264_000.0,
|
||||||
|
10_559_470_521_600.0,
|
||||||
|
670_442_572_800.0,
|
||||||
|
33_522_128_640.0,
|
||||||
|
1_323_241_920.0,
|
||||||
|
40_840_800.0,
|
||||||
|
960_960.0,
|
||||||
|
16_380.0,
|
||||||
|
182.0,
|
||||||
|
1.0,
|
||||||
|
];
|
||||||
|
|
||||||
|
let n = a.nrows();
|
||||||
|
let norm = (0..n)
|
||||||
|
.map(|j| (0..n).map(|i| a[(i, j)].abs()).sum::<f64>())
|
||||||
|
.fold(0.0_f64, f64::max);
|
||||||
|
let s = if norm > THETA_13 {
|
||||||
|
(norm / THETA_13).log2().ceil() as u32
|
||||||
|
} else {
|
||||||
|
0
|
||||||
|
};
|
||||||
|
let a1 = a * (0.5_f64).powi(s as i32);
|
||||||
|
|
||||||
|
let id = DMatrix::<f64>::identity(n, n);
|
||||||
|
let a2 = &a1 * &a1;
|
||||||
|
let a4 = &a2 * &a2;
|
||||||
|
let a6 = &a2 * &a4;
|
||||||
|
|
||||||
|
let u_inner = &a6 * B[13] + &a4 * B[11] + &a2 * B[9];
|
||||||
|
let u_poly = &a6 * &u_inner + &a6 * B[7] + &a4 * B[5] + &a2 * B[3] + &id * B[1];
|
||||||
|
let u = &a1 * &u_poly;
|
||||||
|
let v_inner = &a6 * B[12] + &a4 * B[10] + &a2 * B[8];
|
||||||
|
let v = &a6 * &v_inner + &a6 * B[6] + &a4 * B[4] + &a2 * B[2] + &id * B[0];
|
||||||
|
|
||||||
|
let p = &v + &u;
|
||||||
|
let q = &v - &u;
|
||||||
|
let mut e = q
|
||||||
|
.lu()
|
||||||
|
.solve(&p)
|
||||||
|
.ok_or_else(|| TensorError::numerical("Padé denominator is singular".to_string()))?;
|
||||||
|
for _ in 0..s {
|
||||||
|
e = &e * &e;
|
||||||
|
}
|
||||||
|
Ok(e)
|
||||||
|
}
|
||||||
@@ -0,0 +1,61 @@
|
|||||||
|
//! Tests for the real matrix exponential.
|
||||||
|
|
||||||
|
use crate::{Device, Tensor};
|
||||||
|
|
||||||
|
fn assert_close(a: f32, b: f32, tol: f32, what: &str) {
|
||||||
|
assert!((a - b).abs() < tol, "{what}: {a} vs {b}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn exp_of_zero_is_identity() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
let z = Tensor::from_data(vec![0.0; 9], vec![3, 3], &device).unwrap();
|
||||||
|
let e = z.matrix_exp().unwrap();
|
||||||
|
let v = e.to_vec().unwrap();
|
||||||
|
for i in 0..3 {
|
||||||
|
for j in 0..3 {
|
||||||
|
let expected = if i == j { 1.0 } else { 0.0 };
|
||||||
|
assert_close(v[i * 3 + j], expected, 1e-6, "exp(0)");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn exp_of_nilpotent_matches_closed_form() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
// exp([[0,1],[0,0]]) = [[1,1],[0,1]].
|
||||||
|
let a = Tensor::from_data(vec![0.0, 1.0, 0.0, 0.0], vec![2, 2], &device).unwrap();
|
||||||
|
let v = a.matrix_exp().unwrap().to_vec().unwrap();
|
||||||
|
assert_close(v[0], 1.0, 1e-6, "00");
|
||||||
|
assert_close(v[1], 1.0, 1e-6, "01");
|
||||||
|
assert_close(v[2], 0.0, 1e-6, "10");
|
||||||
|
assert_close(v[3], 1.0, 1e-6, "11");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn exp_of_rotation_generator_is_rotation() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
let theta = 1.2_f32;
|
||||||
|
let a = Tensor::from_data(vec![0.0, -theta, theta, 0.0], vec![2, 2], &device).unwrap();
|
||||||
|
let v = a.matrix_exp().unwrap().to_vec().unwrap();
|
||||||
|
assert_close(v[0], theta.cos(), 1e-5, "cos");
|
||||||
|
assert_close(v[1], -theta.sin(), 1e-5, "−sin");
|
||||||
|
assert_close(v[2], theta.sin(), 1e-5, "sin");
|
||||||
|
assert_close(v[3], theta.cos(), 1e-5, "cos");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn exp_of_diagonal_with_large_entries() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
let a = Tensor::from_data(vec![8.0, 0.0, 0.0, -3.0], vec![2, 2], &device).unwrap();
|
||||||
|
let v = a.matrix_exp().unwrap().to_vec().unwrap();
|
||||||
|
assert!((v[0] - 8.0_f32.exp()).abs() / 8.0_f32.exp() < 1e-4);
|
||||||
|
assert!((v[3] - (-3.0_f32).exp()).abs() < 1e-5);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn exp_rejects_non_square() {
|
||||||
|
let device = Device::cpu();
|
||||||
|
let a = Tensor::from_data(vec![0.0; 6], vec![2, 3], &device).unwrap();
|
||||||
|
assert!(a.matrix_exp().is_err());
|
||||||
|
}
|
||||||
@@ -73,6 +73,7 @@ pub mod decompositions;
|
|||||||
pub mod helpers;
|
pub mod helpers;
|
||||||
pub mod lu;
|
pub mod lu;
|
||||||
pub mod matrix_analysis;
|
pub mod matrix_analysis;
|
||||||
|
pub mod matrix_exp;
|
||||||
pub mod matrix_power;
|
pub mod matrix_power;
|
||||||
pub mod matrix_solve;
|
pub mod matrix_solve;
|
||||||
pub mod types;
|
pub mod types;
|
||||||
@@ -90,6 +91,9 @@ mod simple_test;
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod matrix_power_tests;
|
mod matrix_power_tests;
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod matrix_exp_tests;
|
||||||
|
|
||||||
// Public re-exports for easy access
|
// Public re-exports for easy access
|
||||||
pub use types::{
|
pub use types::{
|
||||||
EigenResult, LUOptions, LUResult, LeastSquaresMethod, LeastSquaresOptions, LeastSquaresResult,
|
EigenResult, LUOptions, LUResult, LeastSquaresMethod, LeastSquaresOptions, LeastSquaresResult,
|
||||||
|
|||||||
@@ -12,9 +12,11 @@ use crate::cusparse::{
|
|||||||
formats::{CuSparseCOO, CuSparseCSR, ToCuSparseIndex, ToCuSparseValue},
|
formats::{CuSparseCOO, CuSparseCSR, ToCuSparseIndex, ToCuSparseValue},
|
||||||
};
|
};
|
||||||
use crate::sparse::{SparseCOO, SparseCSR};
|
use crate::sparse::{SparseCOO, SparseCSR};
|
||||||
use crate::{Device, Result, Shape, Tensor};
|
use crate::{Device, Result, Shape, Tensor, TensorError};
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
use cudarc::driver::safe::CudaContext as CudaDevice;
|
use cudarc::driver::safe::CudaContext as CudaDevice;
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
use cudarc::driver::DevicePtr;
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
|
|
||||||
use super::types::Storage;
|
use super::types::Storage;
|
||||||
use crate::memory::{CompressedStorage, MemoryAccessPattern, MemoryStats};
|
use crate::memory::{CompressedStorage, MemoryAccessPattern, MemoryStats};
|
||||||
use crate::Result;
|
use crate::{Result, TensorError};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
//! FP16/FP32 conversion methods
|
//! FP16/FP32 conversion methods
|
||||||
|
|
||||||
use super::types::Storage;
|
use super::types::{Storage, StorageData, StorageInner};
|
||||||
|
use crate::{DType, Result, TensorError};
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
use half::f16;
|
use half::f16;
|
||||||
|
|||||||
@@ -3,6 +3,12 @@
|
|||||||
use super::types::{
|
use super::types::{
|
||||||
Storage, StorageData,
|
Storage, StorageData,
|
||||||
};
|
};
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
use super::types::{
|
||||||
|
GPU_CPU_BUFFERS, GPU_CPU_BUFFERS_F16, GPU_CPU_BUFFERS_F64,
|
||||||
|
GPU_CPU_MUT_BUFFERS, GPU_CPU_MUT_BUFFERS_F16, GPU_CPU_MUT_BUFFERS_F64,
|
||||||
|
SendPtr,
|
||||||
|
};
|
||||||
|
|
||||||
#[cfg(all(target_os = "macos", feature = "metal"))]
|
#[cfg(all(target_os = "macos", feature = "metal"))]
|
||||||
use objc2_metal::MTLBuffer;
|
use objc2_metal::MTLBuffer;
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
//!
|
//!
|
||||||
//! This module provides operations for combining multiple tensors along specified dimensions.
|
//! This module provides operations for combining multiple tensors along specified dimensions.
|
||||||
|
|
||||||
use crate::{Result, Storage, Tensor, TensorError};
|
use crate::{Device, Result, Storage, Tensor, TensorError};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
impl Tensor {
|
impl Tensor {
|
||||||
|
|||||||
@@ -7,7 +7,7 @@
|
|||||||
//! - `arange`, `full` - Special patterns
|
//! - `arange`, `full` - Special patterns
|
||||||
|
|
||||||
use super::core::Tensor;
|
use super::core::Tensor;
|
||||||
use crate::{DType, Device, DeviceOp, Result, Shape, Storage, TensorError};
|
use crate::{DType, Device, DeviceOp, NodeId, Result, Shape, Storage, TensorError};
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
|
|||||||
@@ -414,6 +414,16 @@ impl Tensor {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Get the CUDA device pointer (`CUdeviceptr`) for use with GPU kernel launches.
|
||||||
|
///
|
||||||
|
/// Returns `Err` when the tensor is not on a CUDA device. The returned
|
||||||
|
/// value is a raw GPU virtual address (a `u64`) that can be passed to
|
||||||
|
/// `cudarc`'s `launch_builder().arg(&ptr)` as a kernel argument.
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
pub fn cuda_device_ptr(&self) -> Result<cudarc::driver::sys::CUdeviceptr> {
|
||||||
|
self.storage.cuda_device_ptr()
|
||||||
|
}
|
||||||
|
|
||||||
/// Get raw data pointer for CUDA kernel access
|
/// Get raw data pointer for CUDA kernel access
|
||||||
///
|
///
|
||||||
/// # Safety
|
/// # Safety
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ pub mod block_manager;
|
|||||||
pub mod pool;
|
pub mod pool;
|
||||||
|
|
||||||
use crate::config::FlashAttentionConfig;
|
use crate::config::FlashAttentionConfig;
|
||||||
|
use crate::{FlashError, FlashResult};
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
use rtx_runtime::CudaBackend;
|
use rtx_runtime::CudaBackend;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|||||||
@@ -504,6 +504,12 @@ impl MambaBlock {
|
|||||||
|
|
||||||
/// Forward pass through `MambaBlock`
|
/// Forward pass through `MambaBlock`
|
||||||
pub fn forward(&self, x: &Tensor) -> Result<MambaOutput> {
|
pub fn forward(&self, x: &Tensor) -> Result<MambaOutput> {
|
||||||
|
// GPU-accelerated path: cuBLAS for projections, CPU loops for SSM scan.
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
if matches!(self.device, Device::Cuda(_)) {
|
||||||
|
return self.forward_cuda(x);
|
||||||
|
}
|
||||||
|
|
||||||
let dims = x.shape().dims().to_vec();
|
let dims = x.shape().dims().to_vec();
|
||||||
let (b, l, d_model) = (dims[0], dims[1], dims[2]);
|
let (b, l, d_model) = (dims[0], dims[1], dims[2]);
|
||||||
let d = self.config.get_d_inner();
|
let d = self.config.get_d_inner();
|
||||||
@@ -624,6 +630,136 @@ impl MambaBlock {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// GPU-accelerated forward pass using cuBLAS for the four linear projections.
|
||||||
|
///
|
||||||
|
/// The two large projections (`in_proj` [b*l,d_model]→[b*l,2d] and `out_proj`
|
||||||
|
/// [b*l,d]→[b*l,d_model]) and the two smaller projections (`x_proj` and
|
||||||
|
/// `dt_proj`) are dispatched to cuBLAS and stay on GPU. The SSM selective
|
||||||
|
/// scan, conv1d, and element-wise activations (SiLU, softplus) remain on CPU
|
||||||
|
/// and require one D2H + one H2D transfer of the intermediate activations.
|
||||||
|
///
|
||||||
|
/// Called automatically by [`Self::forward`] when `self.device` is
|
||||||
|
/// `Device::Cuda(_)` and the `cuda` feature is enabled.
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
fn forward_cuda(&self, x: &Tensor) -> Result<MambaOutput> {
|
||||||
|
let dims = x.shape().dims().to_vec();
|
||||||
|
let (b, l, d_model) = (dims[0], dims[1], dims[2]);
|
||||||
|
let d = self.config.get_d_inner();
|
||||||
|
let n = self.config.d_state;
|
||||||
|
let dt_rank = self.config.get_dt_rank();
|
||||||
|
let kc = self.config.d_conv;
|
||||||
|
let dbc = dt_rank + 2 * n;
|
||||||
|
|
||||||
|
// ── Step 1: in_proj on GPU (cuBLAS) ─────────────────────────────────
|
||||||
|
// in_proj: [d_model, 2*d]. x_flat: [b*l, d_model].
|
||||||
|
// xz = x_flat @ in_proj → [b*l, 2*d].
|
||||||
|
let x_flat = x.view([b * l, d_model])?;
|
||||||
|
let xz_gpu = x_flat.matmul(&self.in_proj)?;
|
||||||
|
|
||||||
|
// D2H once – only [b*l * 2*d] floats.
|
||||||
|
let xz = xz_gpu.to_cpu()?;
|
||||||
|
|
||||||
|
// Split xz into the SSM branch (x_in) and the gate branch (z).
|
||||||
|
let mut x_in = vec![0.0f32; b * l * d];
|
||||||
|
let mut z = vec![0.0f32; b * l * d];
|
||||||
|
for i in 0..b * l {
|
||||||
|
x_in[i * d..i * d + d].copy_from_slice(&xz[i * 2 * d..i * 2 * d + d]);
|
||||||
|
z[i * d..i * d + d].copy_from_slice(&xz[i * 2 * d + d..i * 2 * d + 2 * d]);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Step 2-3: causal conv1d + SiLU on CPU (O(b*l*d*kc), cheap) ──────
|
||||||
|
let conv_w = self.conv1d_weight.to_cpu()?; // [d, 1, kc] stored as [d, kc]
|
||||||
|
let conv_b = match &self.conv1d_bias {
|
||||||
|
Some(t) => Some(t.to_cpu()?),
|
||||||
|
None => None,
|
||||||
|
};
|
||||||
|
let mut u = vec![0.0f32; b * l * d];
|
||||||
|
for bi in 0..b {
|
||||||
|
for li in 0..l {
|
||||||
|
for j in 0..d {
|
||||||
|
let mut acc = conv_b.as_ref().map_or(0.0f32, |cb| cb[j]);
|
||||||
|
for kk in 0..kc {
|
||||||
|
let src = li as isize - (kc as isize - 1) + kk as isize;
|
||||||
|
if src >= 0 {
|
||||||
|
acc += x_in[(bi * l + src as usize) * d + j]
|
||||||
|
* conv_w[j * kc + kk];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
u[(bi * l + li) * d + j] = silu_f32(acc);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Step 4: x_proj on GPU (cuBLAS) ──────────────────────────────────
|
||||||
|
// u_gpu: [b*l, d]. x_proj: [d, dbc].
|
||||||
|
// xdbl = u_gpu @ x_proj → [b*l, dbc].
|
||||||
|
let u_gpu = Tensor::from_vec(u.clone(), &[b * l, d], &self.device)?;
|
||||||
|
let xdbl_gpu = u_gpu.matmul(&self.x_proj)?;
|
||||||
|
|
||||||
|
// D2H xdbl – [b*l * dbc] floats (small: dbc ≈ 48).
|
||||||
|
let xdbl = xdbl_gpu.to_cpu()?;
|
||||||
|
|
||||||
|
// ── Step 5: extract dt, B, C; dt_proj + softplus on CPU ─────────────
|
||||||
|
let dt_proj_w = self.dt_proj.to_cpu()?; // [dt_rank, d]
|
||||||
|
let dt_bias = self.dt_bias.to_cpu()?; // [d]
|
||||||
|
let a_log = self.A_log.to_cpu()?; // [d, n]
|
||||||
|
let d_skip = self.d_skip.to_cpu()?; // [d]
|
||||||
|
|
||||||
|
let mut bmat = vec![0.0f32; b * l * n];
|
||||||
|
let mut cmat = vec![0.0f32; b * l * n];
|
||||||
|
let mut delta = vec![0.0f32; b * l * d];
|
||||||
|
for i in 0..b * l {
|
||||||
|
for nn in 0..n {
|
||||||
|
bmat[i * n + nn] = xdbl[i * dbc + dt_rank + nn];
|
||||||
|
cmat[i * n + nn] = xdbl[i * dbc + dt_rank + n + nn];
|
||||||
|
}
|
||||||
|
for j in 0..d {
|
||||||
|
let mut s = dt_bias[j];
|
||||||
|
for r in 0..dt_rank {
|
||||||
|
s += xdbl[i * dbc + r] * dt_proj_w[r * d + j];
|
||||||
|
}
|
||||||
|
delta[i * d + j] = softplus_f32(s);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Step 6-7: selective scan + SiLU gate on CPU ──────────────────────
|
||||||
|
let mut out_flat = vec![0.0f32; b * l * d];
|
||||||
|
for bi in 0..b {
|
||||||
|
let mut h = vec![0.0f32; d * n];
|
||||||
|
for li in 0..l {
|
||||||
|
let base = (bi * l + li) * d;
|
||||||
|
for j in 0..d {
|
||||||
|
let dj = delta[base + j];
|
||||||
|
let uj = u[base + j];
|
||||||
|
let mut yj = d_skip[j] * uj;
|
||||||
|
for nn in 0..n {
|
||||||
|
let a = -a_log[j * n + nn].exp();
|
||||||
|
let da = (dj * a).exp();
|
||||||
|
let dbu = dj * bmat[(bi * l + li) * n + nn] * uj;
|
||||||
|
let hv = da * h[j * n + nn] + dbu;
|
||||||
|
h[j * n + nn] = hv;
|
||||||
|
yj += cmat[(bi * l + li) * n + nn] * hv;
|
||||||
|
}
|
||||||
|
// SiLU gate fused into the scan output.
|
||||||
|
out_flat[base + j] = yj * silu_f32(z[base + j]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Step 8: out_proj on GPU (cuBLAS) ─────────────────────────────────
|
||||||
|
// H2D once – [b*l * d] floats.
|
||||||
|
let y_gated_gpu = Tensor::from_vec(out_flat, &[b * l, d], &self.device)?;
|
||||||
|
// out_proj: [d, d_model].
|
||||||
|
let out_flat_gpu = y_gated_gpu.matmul(&self.out_proj)?;
|
||||||
|
|
||||||
|
// Reshape to [b, l, d_model] — zero-copy view.
|
||||||
|
let output = out_flat_gpu.view([b, l, d_model])?;
|
||||||
|
Ok(MambaOutput {
|
||||||
|
output,
|
||||||
|
aux_info: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
/// Analytic backward pass: given `d_out = ∂L/∂out` (same shape as
|
/// Analytic backward pass: given `d_out = ∂L/∂out` (same shape as
|
||||||
/// the forward output, `[b,l,d_model]`), return `∂L/∂θ` for every
|
/// the forward output, `[b,l,d_model]`), return `∂L/∂θ` for every
|
||||||
/// parameter, keyed by its persistence name (`in_proj`,
|
/// parameter, keyed by its persistence name (`in_proj`,
|
||||||
|
|||||||
@@ -19,9 +19,13 @@ use rtx_tensor::{Tensor, Device, DType};
|
|||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
|
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
use cudarc::driver::{CudaDevice, LaunchAsync, LaunchConfig};
|
use cudarc::driver::{CudaContext, CudaFunction, CudaStream, LaunchConfig, PushKernelArg};
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
use cudarc::driver::sys::CUdeviceptr;
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
use cudarc::nvrtc::compile_ptx;
|
use cudarc::nvrtc::compile_ptx;
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
/// CUDA kernel source code for forward selective scan
|
/// CUDA kernel source code for forward selective scan
|
||||||
const SELECTIVE_SCAN_FORWARD_KERNEL: &str = r#"
|
const SELECTIVE_SCAN_FORWARD_KERNEL: &str = r#"
|
||||||
@@ -258,11 +262,8 @@ impl Default for KernelConfig {
|
|||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct MambaCudaKernels {
|
pub struct MambaCudaKernels {
|
||||||
/// CUDA device handle
|
stream: Arc<CudaStream>,
|
||||||
device: CudaDevice,
|
kernels: HashMap<String, CudaFunction>,
|
||||||
/// Compiled kernel modules
|
|
||||||
kernels: HashMap<String, cudarc::driver::CudaFunction>,
|
|
||||||
/// Kernel configuration
|
|
||||||
config: KernelConfig,
|
config: KernelConfig,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -276,39 +277,35 @@ pub struct MambaCudaKernels {
|
|||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
impl MambaCudaKernels {
|
impl MambaCudaKernels {
|
||||||
/// Create new CUDA kernel manager
|
/// Create new CUDA kernel manager
|
||||||
pub fn new(device: CudaDevice, config: KernelConfig) -> Result<Self> {
|
pub fn new(ctx: Arc<CudaContext>, config: KernelConfig) -> Result<Self> {
|
||||||
|
let stream = ctx.default_stream();
|
||||||
let mut kernels = HashMap::new();
|
let mut kernels = HashMap::new();
|
||||||
|
|
||||||
// Compile kernels
|
|
||||||
let ptx = compile_ptx(SELECTIVE_SCAN_FORWARD_KERNEL)
|
let ptx = compile_ptx(SELECTIVE_SCAN_FORWARD_KERNEL)
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to compile forward kernel: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to compile forward kernel: {}", e)))?;
|
||||||
let module = device.load_ptx(ptx, "selective_scan_forward", &["selective_scan_forward"])
|
let module = ctx.load_module(ptx)
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to load forward kernel: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to load forward kernel: {}", e)))?;
|
||||||
let func = module.get_func("selective_scan_forward")
|
let func = module.load_function("selective_scan_forward")
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to get forward kernel function: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to get forward kernel function: {}", e)))?;
|
||||||
kernels.insert("selective_scan_forward".to_string(), func);
|
kernels.insert("selective_scan_forward".to_string(), func);
|
||||||
|
|
||||||
let ptx = compile_ptx(SELECTIVE_SCAN_BACKWARD_KERNEL)
|
let ptx = compile_ptx(SELECTIVE_SCAN_BACKWARD_KERNEL)
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to compile backward kernel: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to compile backward kernel: {}", e)))?;
|
||||||
let module = device.load_ptx(ptx, "selective_scan_backward", &["selective_scan_backward"])
|
let module = ctx.load_module(ptx)
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to load backward kernel: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to load backward kernel: {}", e)))?;
|
||||||
let func = module.get_func("selective_scan_backward")
|
let func = module.load_function("selective_scan_backward")
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to get backward kernel function: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to get backward kernel function: {}", e)))?;
|
||||||
kernels.insert("selective_scan_backward".to_string(), func);
|
kernels.insert("selective_scan_backward".to_string(), func);
|
||||||
|
|
||||||
let ptx = compile_ptx(CAUSAL_CONV1D_KERNEL)
|
let ptx = compile_ptx(CAUSAL_CONV1D_KERNEL)
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to compile conv1d kernel: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to compile conv1d kernel: {}", e)))?;
|
||||||
let module = device.load_ptx(ptx, "causal_conv1d", &["causal_conv1d_forward"])
|
let module = ctx.load_module(ptx)
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to load conv1d kernel: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to load conv1d kernel: {}", e)))?;
|
||||||
let func = module.get_func("causal_conv1d_forward")
|
let func = module.load_function("causal_conv1d_forward")
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to get conv1d kernel function: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to get conv1d kernel function: {}", e)))?;
|
||||||
kernels.insert("causal_conv1d_forward".to_string(), func);
|
kernels.insert("causal_conv1d_forward".to_string(), func);
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self { stream, kernels, config })
|
||||||
device,
|
|
||||||
kernels,
|
|
||||||
config,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Launch forward selective scan kernel
|
/// Launch forward selective scan kernel
|
||||||
@@ -343,28 +340,44 @@ impl MambaCudaKernels {
|
|||||||
shared_mem_bytes,
|
shared_mem_bytes,
|
||||||
};
|
};
|
||||||
|
|
||||||
// For now, use placeholder pointers until rtx-tensor provides GPU memory access
|
let u_ptr: CUdeviceptr = u.cuda_device_ptr()
|
||||||
let params = (
|
.map_err(|e| TransformerError::CudaRuntime(format!("u not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // u.as_ptr(),
|
let delta_ptr: CUdeviceptr = delta.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // delta.as_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("delta not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // A.as_ptr(),
|
let a_ptr: CUdeviceptr = A.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // B.as_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("A not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // C.as_ptr(),
|
let b_ptr: CUdeviceptr = B.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // D.map(|d| d.as_ptr()).unwrap_or(std::ptr::null()),
|
.map_err(|e| TransformerError::CudaRuntime(format!("B not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null_mut::<f32>(), // output.as_mut_ptr(),
|
let c_ptr: CUdeviceptr = C.cuda_device_ptr()
|
||||||
batch_size,
|
.map_err(|e| TransformerError::CudaRuntime(format!("C not on CUDA device: {}", e)))?;
|
||||||
seq_len,
|
let d_ptr: CUdeviceptr = D
|
||||||
d_model,
|
.map(|d| d.cuda_device_ptr())
|
||||||
d_state,
|
.transpose()
|
||||||
);
|
.map_err(|e| TransformerError::CudaRuntime(format!("D not on CUDA device: {}", e)))?
|
||||||
|
.unwrap_or(0);
|
||||||
|
let out_ptr: CUdeviceptr = output.cuda_device_ptr()
|
||||||
|
.map_err(|e| TransformerError::CudaRuntime(format!("output not on CUDA device: {}", e)))?;
|
||||||
|
|
||||||
unsafe {
|
unsafe {
|
||||||
kernel.launch(launch_config, params)
|
self.stream
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to launch kernel: {}", e)))?;
|
.launch_builder(kernel)
|
||||||
|
.arg(&u_ptr)
|
||||||
|
.arg(&delta_ptr)
|
||||||
|
.arg(&a_ptr)
|
||||||
|
.arg(&b_ptr)
|
||||||
|
.arg(&c_ptr)
|
||||||
|
.arg(&d_ptr)
|
||||||
|
.arg(&out_ptr)
|
||||||
|
.arg(&batch_size)
|
||||||
|
.arg(&seq_len)
|
||||||
|
.arg(&d_model)
|
||||||
|
.arg(&d_state)
|
||||||
|
.launch(launch_config)
|
||||||
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to launch selective_scan_forward: {}", e)))?;
|
||||||
}
|
}
|
||||||
|
|
||||||
self.device.synchronize()
|
self.stream.synchronize()
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to synchronize device: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to synchronize stream: {}", e)))?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -402,32 +415,53 @@ impl MambaCudaKernels {
|
|||||||
shared_mem_bytes,
|
shared_mem_bytes,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Placeholder pointers until rtx-tensor provides GPU memory access
|
let grad_out_ptr: CUdeviceptr = grad_output.cuda_device_ptr()
|
||||||
let params = (
|
.map_err(|e| TransformerError::CudaRuntime(format!("grad_output not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // grad_output.as_ptr(),
|
let u_ptr: CUdeviceptr = u.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // u.as_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("u not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // delta.as_ptr(),
|
let delta_ptr: CUdeviceptr = delta.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // A.as_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("delta not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // B.as_ptr(),
|
let a_ptr: CUdeviceptr = A.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // C.as_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("A not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null_mut::<f32>(), // grad_u.as_mut_ptr(),
|
let b_ptr: CUdeviceptr = B.cuda_device_ptr()
|
||||||
std::ptr::null_mut::<f32>(), // grad_delta.as_mut_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("B not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null_mut::<f32>(), // grad_A.as_mut_ptr(),
|
let c_ptr: CUdeviceptr = C.cuda_device_ptr()
|
||||||
std::ptr::null_mut::<f32>(), // grad_B.as_mut_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("C not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null_mut::<f32>(), // grad_C.as_mut_ptr(),
|
let grad_u_ptr: CUdeviceptr = grad_u.cuda_device_ptr()
|
||||||
batch_size,
|
.map_err(|e| TransformerError::CudaRuntime(format!("grad_u not on CUDA device: {}", e)))?;
|
||||||
seq_len,
|
let grad_delta_ptr: CUdeviceptr = grad_delta.cuda_device_ptr()
|
||||||
d_model,
|
.map_err(|e| TransformerError::CudaRuntime(format!("grad_delta not on CUDA device: {}", e)))?;
|
||||||
d_state,
|
let grad_a_ptr: CUdeviceptr = grad_A.cuda_device_ptr()
|
||||||
);
|
.map_err(|e| TransformerError::CudaRuntime(format!("grad_A not on CUDA device: {}", e)))?;
|
||||||
|
let grad_b_ptr: CUdeviceptr = grad_B.cuda_device_ptr()
|
||||||
|
.map_err(|e| TransformerError::CudaRuntime(format!("grad_B not on CUDA device: {}", e)))?;
|
||||||
|
let grad_c_ptr: CUdeviceptr = grad_C.cuda_device_ptr()
|
||||||
|
.map_err(|e| TransformerError::CudaRuntime(format!("grad_C not on CUDA device: {}", e)))?;
|
||||||
|
|
||||||
unsafe {
|
unsafe {
|
||||||
kernel.launch(launch_config, params)
|
self.stream
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to launch kernel: {}", e)))?;
|
.launch_builder(kernel)
|
||||||
|
.arg(&grad_out_ptr)
|
||||||
|
.arg(&u_ptr)
|
||||||
|
.arg(&delta_ptr)
|
||||||
|
.arg(&a_ptr)
|
||||||
|
.arg(&b_ptr)
|
||||||
|
.arg(&c_ptr)
|
||||||
|
.arg(&grad_u_ptr)
|
||||||
|
.arg(&grad_delta_ptr)
|
||||||
|
.arg(&grad_a_ptr)
|
||||||
|
.arg(&grad_b_ptr)
|
||||||
|
.arg(&grad_c_ptr)
|
||||||
|
.arg(&batch_size)
|
||||||
|
.arg(&seq_len)
|
||||||
|
.arg(&d_model)
|
||||||
|
.arg(&d_state)
|
||||||
|
.launch(launch_config)
|
||||||
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to launch selective_scan_backward: {}", e)))?;
|
||||||
}
|
}
|
||||||
|
|
||||||
self.device.synchronize()
|
self.stream.synchronize()
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to synchronize device: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to synchronize stream: {}", e)))?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -458,25 +492,35 @@ impl MambaCudaKernels {
|
|||||||
shared_mem_bytes: 0,
|
shared_mem_bytes: 0,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Placeholder pointers until rtx-tensor provides GPU memory access
|
let input_ptr: CUdeviceptr = input.cuda_device_ptr()
|
||||||
let params = (
|
.map_err(|e| TransformerError::CudaRuntime(format!("input not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // input.as_ptr(),
|
let weight_ptr: CUdeviceptr = weight.cuda_device_ptr()
|
||||||
std::ptr::null::<f32>(), // weight.as_ptr(),
|
.map_err(|e| TransformerError::CudaRuntime(format!("weight not on CUDA device: {}", e)))?;
|
||||||
std::ptr::null::<f32>(), // bias.map(|b| b.as_ptr()).unwrap_or(std::ptr::null()),
|
let bias_ptr: CUdeviceptr = bias
|
||||||
std::ptr::null_mut::<f32>(), // output.as_mut_ptr(),
|
.map(|b| b.cuda_device_ptr())
|
||||||
batch_size,
|
.transpose()
|
||||||
d_model,
|
.map_err(|e| TransformerError::CudaRuntime(format!("bias not on CUDA device: {}", e)))?
|
||||||
seq_len,
|
.unwrap_or(0);
|
||||||
kernel_size,
|
let out_ptr: CUdeviceptr = output.cuda_device_ptr()
|
||||||
);
|
.map_err(|e| TransformerError::CudaRuntime(format!("output not on CUDA device: {}", e)))?;
|
||||||
|
|
||||||
unsafe {
|
unsafe {
|
||||||
kernel.launch(launch_config, params)
|
self.stream
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to launch kernel: {}", e)))?;
|
.launch_builder(kernel)
|
||||||
|
.arg(&input_ptr)
|
||||||
|
.arg(&weight_ptr)
|
||||||
|
.arg(&bias_ptr)
|
||||||
|
.arg(&out_ptr)
|
||||||
|
.arg(&batch_size)
|
||||||
|
.arg(&d_model)
|
||||||
|
.arg(&seq_len)
|
||||||
|
.arg(&kernel_size)
|
||||||
|
.launch(launch_config)
|
||||||
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to launch causal_conv1d_forward: {}", e)))?;
|
||||||
}
|
}
|
||||||
|
|
||||||
self.device.synchronize()
|
self.stream.synchronize()
|
||||||
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to synchronize device: {}", e)))?;
|
.map_err(|e| TransformerError::CudaRuntime(format!("Failed to synchronize stream: {}", e)))?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user