Files
rustytorch/crates/specialized/rtx-compiler/src/lib.rs
T
2026-03-04 00:08:42 +00:00

877 lines
28 KiB
Rust

//! RTX Compiler - GPU kernel compilation via rustg
//!
//! This crate provides the interface to the rustg compiler for
//! compiling GPU kernels targeting RTX 5090 (sm_120) and other architectures.
use anyhow::Result;
use std::collections::HashMap;
use std::fmt::Write as FmtWrite;
use std::hash::{Hash, Hasher};
use std::path::PathBuf;
use std::time::Duration;
use thiserror::Error;
use tracing::{debug, info, warn};
pub mod codegen;
pub mod ir;
pub mod optimizer;
mod rustg_backend;
/// Compiler errors
#[derive(Debug, Error)]
pub enum CompilerError {
/// rustg compilation failed
#[error("rustg compilation failed: {0}")]
RustgCompilationFailed(String),
/// rustg binary not found
#[error("rustg binary not found: {0}")]
RustgNotFound(String),
/// Compilation timeout
#[error("Compilation timeout after {0} seconds")]
CompilationTimeout(u64),
/// Invalid kernel configuration
#[error("Invalid kernel configuration: {0}")]
InvalidKernelConfig(String),
/// PTX validation failed
#[error("PTX validation failed: {0}")]
PtxValidationFailed(String),
/// Cache operation failed
#[error("Kernel cache operation failed: {0}")]
CacheError(String),
/// Unsupported target architecture
#[error("Unsupported target architecture: {0:?}")]
UnsupportedTarget(Target),
/// Source hash computation failed
#[error("Source hash computation failed: {0}")]
HashError(String),
/// IO error
#[error("IO error: {0}")]
Io(#[from] std::io::Error),
}
/// Represents a compilation target architecture
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
pub enum Target {
/// NVIDIA RTX 5090 (sm_120)
SM120,
/// NVIDIA H100 (sm_90)
SM90,
/// AMD MI300 (gfx942)
GFX942,
}
impl Target {
/// Get architecture string for rustg
pub fn as_str(&self) -> &str {
match self {
Self::SM120 => "sm_120",
Self::SM90 => "sm_90",
Self::GFX942 => "gfx942",
}
}
}
/// Compilation options
#[derive(Debug, Clone)]
pub struct CompileOptions {
/// Target GPU architecture
pub target: Target,
/// Enable optimizations
pub optimize: bool,
/// Enable debug info
pub debug_info: bool,
/// Cache directory for compiled kernels
pub cache_dir: PathBuf,
/// Maximum registers per thread
pub max_registers: Option<u32>,
/// Enable fast math optimizations
pub fast_math: bool,
/// Dead code elimination
pub dead_code_elimination: bool,
/// Constant folding
pub constant_folding: bool,
/// Loop unrolling
pub loop_unrolling: bool,
/// Shared memory banking optimization
pub shared_memory_banking: bool,
/// Register spilling threshold
pub register_spill_threshold: Option<f32>,
/// Enable tensor core optimizations (RTX 5090 specific)
pub tensor_core_optimizations: bool,
/// Compilation timeout in seconds
pub timeout_seconds: u64,
/// Additional rustg arguments
pub additional_args: Vec<String>,
}
impl Default for CompileOptions {
fn default() -> Self {
Self {
target: Target::SM120,
optimize: true,
debug_info: false,
cache_dir: PathBuf::from("target/kernel_cache"),
max_registers: None,
fast_math: false,
dead_code_elimination: true,
constant_folding: true,
loop_unrolling: true,
shared_memory_banking: true,
register_spill_threshold: Some(0.8),
tensor_core_optimizations: true,
timeout_seconds: 60,
additional_args: Vec::new(),
}
}
}
/// Compilation metrics for performance tracking
#[derive(Debug, Clone, Default)]
pub struct CompilationMetrics {
/// Total compilation time
pub compilation_time: Duration,
/// Generated PTX size in bytes
pub generated_code_size: usize,
/// Number of optimization passes applied
pub optimization_passes: u32,
/// Register usage estimate
pub estimated_registers: u32,
/// Shared memory usage in bytes
pub shared_memory_usage: u32,
/// Source code hash for cache validation
pub source_hash: u64,
/// Cache hit/miss status
pub cache_hit: bool,
/// rustg version used
pub rustg_version: Option<String>,
}
/// Main compiler struct
pub struct RtxCompiler {
options: CompileOptions,
metrics: HashMap<String, CompilationMetrics>,
}
impl RtxCompiler {
/// Create new compiler with options
pub fn new(options: CompileOptions) -> Self {
Self {
options,
metrics: HashMap::new(),
}
}
/// Create compiler with default options
pub fn default() -> Self {
Self::new(CompileOptions::default())
}
/// Get compilation metrics for a kernel
pub fn get_metrics(&self, kernel_name: &str) -> Option<&CompilationMetrics> {
self.metrics.get(kernel_name)
}
/// Get all compilation metrics
pub fn get_all_metrics(&self) -> &HashMap<String, CompilationMetrics> {
&self.metrics
}
/// Get cache path for kernel
fn get_cache_path(&self, kernel_name: &str) -> PathBuf {
let mut path = self.options.cache_dir.clone();
path.push(self.options.target.as_str());
path.push(format!("{kernel_name}.ptx"));
path
}
/// Get cache metadata path for kernel
fn get_cache_metadata_path(&self, kernel_name: &str) -> PathBuf {
let mut path = self.options.cache_dir.clone();
path.push(self.options.target.as_str());
path.push(format!("{kernel_name}.meta"));
path
}
/// Compute source hash for cache validation
fn compute_source_hash(&self, source: &str) -> u64 {
use std::collections::hash_map::DefaultHasher;
let mut hasher = DefaultHasher::new();
source.hash(&mut hasher);
// Hash the options manually since f32 doesn't implement Hash
self.options.target.hash(&mut hasher);
self.options.optimize.hash(&mut hasher);
self.options.debug_info.hash(&mut hasher);
self.options.max_registers.hash(&mut hasher);
self.options.fast_math.hash(&mut hasher);
self.options.dead_code_elimination.hash(&mut hasher);
self.options.constant_folding.hash(&mut hasher);
self.options.loop_unrolling.hash(&mut hasher);
self.options.shared_memory_banking.hash(&mut hasher);
self.options.tensor_core_optimizations.hash(&mut hasher);
self.options.timeout_seconds.hash(&mut hasher);
// Skip register_spill_threshold (f32) and additional_args for now
hasher.finish()
}
/// Check kernel cache with source validation
fn check_cache(
&self,
kernel_name: &str,
source: &str,
) -> Result<Option<Vec<u8>>, CompilerError> {
let cache_path = self.get_cache_path(kernel_name);
let metadata_path = self.get_cache_metadata_path(kernel_name);
if !cache_path.exists() || !metadata_path.exists() {
return Ok(None);
}
// Check if source has changed by comparing hashes
let current_hash = self.compute_source_hash(source);
match std::fs::read_to_string(&metadata_path) {
Ok(metadata_content) => {
let lines: Vec<&str> = metadata_content.lines().collect();
if !lines.is_empty()
&& let Ok(cached_hash) = lines[0].parse::<u64>()
{
if cached_hash == current_hash {
debug!("Cache hit for kernel '{}'", kernel_name);
return std::fs::read(&cache_path)
.map(Some)
.map_err(CompilerError::Io);
}
debug!("Cache miss for kernel '{}' - source changed", kernel_name);
}
}
Err(_) => {
debug!("Cache metadata invalid for kernel '{}'", kernel_name);
}
}
Ok(None)
}
/// Cache compiled kernel with metadata
fn cache_kernel(
&self,
kernel_name: &str,
code: &[u8],
source: &str,
metrics: &CompilationMetrics,
) -> Result<(), CompilerError> {
let cache_path = self.get_cache_path(kernel_name);
let metadata_path = self.get_cache_metadata_path(kernel_name);
// Ensure cache directory exists
if let Some(parent) = cache_path.parent() {
std::fs::create_dir_all(parent).map_err(CompilerError::Io)?;
}
// Write the compiled code
std::fs::write(&cache_path, code).map_err(CompilerError::Io)?;
// Write metadata
let metadata_content = format!(
"{}\n{}\n{}\n{}\n{}",
self.compute_source_hash(source),
metrics.compilation_time.as_millis(),
metrics.generated_code_size,
metrics.estimated_registers,
metrics.shared_memory_usage
);
std::fs::write(&metadata_path, metadata_content).map_err(CompilerError::Io)?;
debug!("Cached kernel '{}' at {:?}", kernel_name, cache_path);
Ok(())
}
/// Compile Rust GPU code to PTX/SASS
pub fn compile_kernel(
&mut self,
kernel_name: &str,
source: &str,
) -> Result<Vec<u8>, CompilerError> {
info!(
"Compiling kernel '{}' for {}",
kernel_name,
self.options.target.as_str()
);
// Check cache first
if let Some(cached) = self.check_cache(kernel_name, source)? {
debug!("Using cached kernel for '{}'", kernel_name);
// Update metrics for cache hit
let mut metrics = CompilationMetrics::default();
metrics.cache_hit = true;
metrics.source_hash = self.compute_source_hash(source);
metrics.generated_code_size = cached.len();
self.metrics.insert(kernel_name.to_string(), metrics);
return Ok(cached);
}
// Compile with rustg
let (compiled, mut metrics) = self.compile_with_rustg(source)?;
metrics.cache_hit = false;
// Cache the result
self.cache_kernel(kernel_name, &compiled, source, &metrics)?;
// Store metrics
self.metrics.insert(kernel_name.to_string(), metrics);
Ok(compiled)
}
/// Compile with rustg backend
fn compile_with_rustg(
&self,
source: &str,
) -> Result<(Vec<u8>, CompilationMetrics), CompilerError> {
let (ptx, mut metrics) =
rustg_backend::RustgBackend::compile_with_rustg(source, &self.options)?;
metrics.source_hash = self.compute_source_hash(source);
Ok((ptx, metrics))
}
/// Validate compilation result
pub fn validate_compilation(&self, kernel_name: &str, ptx: &[u8]) -> Result<(), CompilerError> {
rustg_backend::RustgBackend::validate_ptx_content(ptx, self.options.target)?;
// Additional validation specific to kernel
let ptx_str = String::from_utf8_lossy(ptx);
// Check that the kernel name appears in PTX
if !ptx_str.contains(kernel_name) && !ptx_str.contains("kernel") {
warn!("Kernel name '{}' not found in generated PTX", kernel_name);
}
// Validate architecture-specific features
if self.options.target == Target::SM120
&& self.options.tensor_core_optimizations
&& self.options.optimize
{
// For RTX 5090, we should see some optimized patterns
// This is a simplified check
if ptx_str.len() < 200 {
warn!("RTX 5090 kernel seems too simple for tensor core optimizations");
}
}
Ok(())
}
/// Profile compilation performance
pub fn profile_compilation(&self, kernel_name: &str) -> Option<String> {
if let Some(metrics) = self.get_metrics(kernel_name) {
let mut profile = String::new();
let _ = writeln!(profile, "Compilation Profile for '{kernel_name}'");
let _ = writeln!(
profile,
" Compilation time: {:?}",
metrics.compilation_time
);
let _ = writeln!(
profile,
" Generated code size: {} bytes",
metrics.generated_code_size
);
let _ = writeln!(
profile,
" Estimated registers: {}",
metrics.estimated_registers
);
let _ = writeln!(
profile,
" Shared memory usage: {} bytes",
metrics.shared_memory_usage
);
let _ = writeln!(profile, " Cache hit: {}", metrics.cache_hit);
let _ = writeln!(
profile,
" Optimization passes: {}",
metrics.optimization_passes
);
if let Some(ref version) = metrics.rustg_version {
let _ = writeln!(profile, " Rustg version: {version}");
}
Some(profile)
} else {
None
}
}
}
/// Compiles a kernel for the specified target
///
/// # Arguments
/// * `name` - Name of the kernel
/// * `source` - Kernel source code
/// * `target` - Target architecture
///
/// # Returns
/// Compiled kernel binary or error
pub fn compile_kernel(name: &str, source: &str, target: Target) -> Result<Vec<u8>> {
let options = CompileOptions {
target,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
compiler
.compile_kernel(name, source)
.map_err(|e| anyhow::anyhow!(e))
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
// Unused import removed
const SIMPLE_KERNEL_SOURCE: &str = r#"
#![no_std]
use rustg_cuda::prelude::*;
#[kernel]
pub fn vector_add(a: &[f32], b: &[f32], c: &mut [f32]) {
let idx = thread_idx_x() + block_idx_x() * block_dim_x();
if idx < c.len() {
c[idx] = a[idx] + b[idx];
}
}
"#;
const RTX5090_KERNEL_SOURCE: &str = r#"
#![no_std]
use rustg_cuda::prelude::*;
#[kernel]
pub fn rtx5090_tensor_core_gemm(
a: &TensorCoreMatrix<f16>,
b: &TensorCoreMatrix<f16>,
c: &mut TensorCoreMatrix<f32>
) {
let warp_id = thread_idx_x() / 32;
let lane_id = thread_idx_x() % 32;
// Use 4th generation Tensor Cores
tensor_core_mma_sm120(a, b, c, warp_id, lane_id);
}
"#;
#[test]
fn test_compile_kernel_basic() {
let result = compile_kernel("test", "kernel code", Target::SM120);
assert!(result.is_ok());
assert!(!result.unwrap().is_empty());
}
#[test]
fn test_rustg_integration_simple_kernel() {
let options = CompileOptions {
target: Target::SM120,
optimize: true,
debug_info: false,
cache_dir: std::env::temp_dir().join("rtx_test_cache"),
max_registers: None,
fast_math: false,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
let result = compiler.compile_kernel("vector_add", SIMPLE_KERNEL_SOURCE);
assert!(result.is_ok(), "Simple kernel compilation should succeed");
let ptx = result.unwrap();
let ptx_str = String::from_utf8_lossy(&ptx);
// Verify PTX structure for RTX 5090 (sm_120)
assert!(
ptx_str.contains(".version"),
"PTX should contain version directive"
);
assert!(
ptx_str.contains(".target sm_120"),
"PTX should target sm_120 for RTX 5090"
);
assert!(
ptx_str.contains(".address_size 64"),
"PTX should use 64-bit addressing"
);
assert!(
ptx_str.contains(".visible .entry"),
"PTX should contain kernel entry point"
);
// Should not contain placeholder comments
assert!(
!ptx_str.contains("// Source:"),
"Should not contain placeholder PTX"
);
}
#[test]
fn test_rtx5090_tensor_core_compilation() {
let options = CompileOptions {
target: Target::SM120,
optimize: true,
debug_info: false,
cache_dir: std::env::temp_dir().join("rtx_test_cache_tensor"),
max_registers: Some(128),
fast_math: true,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
let result = compiler.compile_kernel("tensor_core_gemm", RTX5090_KERNEL_SOURCE);
assert!(result.is_ok(), "RTX 5090 tensor core kernel should compile");
let ptx = result.unwrap();
let ptx_str = String::from_utf8_lossy(&ptx);
// Verify RTX 5090 specific features
assert!(
ptx_str.contains(".target sm_120"),
"Should target RTX 5090 architecture"
);
assert!(
ptx_str.contains(".maxnreg 128"),
"Should respect register limit"
);
// Check for tensor core instructions (these would be generated by rustg)
// Note: Real rustg would generate mma.sync instructions for tensor cores
assert!(
ptx_str.len() > 500,
"RTX 5090 kernel should generate substantial PTX code"
);
}
#[test]
fn test_compilation_error_handling() {
let mut compiler = RtxCompiler::default();
// Test invalid source code
let invalid_source = "This is not valid Rust code!";
let result = compiler.compile_kernel("invalid", invalid_source);
assert!(result.is_err(), "Invalid source should fail compilation");
match result.unwrap_err() {
CompilerError::RustgCompilationFailed(msg) => {
assert!(
msg.contains("syntax error") || msg.contains("parse error"),
"Error message should indicate syntax error: {}",
msg
);
}
_ => panic!("Expected RustgCompilationFailed error"),
}
}
#[test]
fn test_cross_architecture_compilation() {
let targets = [Target::SM120, Target::SM90, Target::GFX942];
for target in &targets {
let options = CompileOptions {
target: *target,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
let result = compiler.compile_kernel("cross_arch_test", SIMPLE_KERNEL_SOURCE);
match target {
Target::SM120 | Target::SM90 => {
assert!(result.is_ok(), "NVIDIA targets should compile successfully");
let ptx = result.unwrap();
let ptx_str = String::from_utf8_lossy(&ptx);
assert!(
ptx_str.contains(&format!(".target {}", target.as_str())),
"PTX should target correct architecture"
);
}
Target::GFX942 => {
// AMD compilation might require different handling
if result.is_ok() {
let code = result.unwrap();
assert!(!code.is_empty(), "AMD compilation should produce code");
}
}
}
}
}
#[test]
fn test_optimization_passes_applied() {
let options = CompileOptions {
target: Target::SM120,
optimize: true,
fast_math: true,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
let optimized_source = r#"
#![no_std]
use rustg_cuda::prelude::*;
#[kernel]
pub fn optimizable_kernel(data: &mut [f32]) {
let idx = thread_idx_x();
let dead_var = 42.0; // Should be eliminated
let const_expr = 2.0 * 3.0; // Should be folded to 6.0
if idx < data.len() {
data[idx] = data[idx] * const_expr + 1.0;
}
}
"#;
let result = compiler.compile_kernel("optimizable", optimized_source);
assert!(result.is_ok(), "Optimizable kernel should compile");
let ptx = result.unwrap();
let ptx_str = String::from_utf8_lossy(&ptx);
// Check for optimization indicators
// Real rustg would apply dead code elimination and constant folding
assert!(
!ptx_str.contains("dead_var"),
"Dead variable should be eliminated"
);
assert!(
ptx_str.len() > 200,
"Optimized kernel should still produce substantial code"
);
}
#[test]
fn test_cache_functionality() {
let cache_dir = std::env::temp_dir().join("rtx_cache_test");
let _ = fs::remove_dir_all(&cache_dir); // Clean up any previous test
let options = CompileOptions {
target: Target::SM120,
cache_dir: cache_dir.clone(),
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
// First compilation
let result1 = compiler.compile_kernel("cached_kernel", SIMPLE_KERNEL_SOURCE);
assert!(result1.is_ok(), "First compilation should succeed");
let cache_path = cache_dir.join("sm_120").join("cached_kernel.ptx");
assert!(cache_path.exists(), "Cache file should be created");
// Second compilation should use cache
let result2 = compiler.compile_kernel("cached_kernel", SIMPLE_KERNEL_SOURCE);
assert!(result2.is_ok(), "Cached compilation should succeed");
// Results should be identical
assert_eq!(
result1.unwrap(),
result2.unwrap(),
"Cached result should match original"
);
// Clean up
let _ = fs::remove_dir_all(&cache_dir);
}
#[test]
fn test_compilation_metrics_tracking() {
let mut compiler = RtxCompiler::default();
// This test will verify that compilation metrics are tracked
// (compilation time, generated code size, optimization passes applied, etc.)
let start_time = std::time::Instant::now();
let result = compiler.compile_kernel("metrics_test", SIMPLE_KERNEL_SOURCE);
let _compilation_time = start_time.elapsed();
assert!(result.is_ok(), "Metrics test compilation should succeed");
let ptx = result.unwrap();
// Verify that we have meaningful compilation output
assert!(ptx.len() > 100, "Generated PTX should be substantial");
// In the real implementation, compilation metrics would be available
// through the compiler interface
}
#[test]
fn test_shared_memory_optimization() {
let shared_memory_kernel = r#"
#![no_std]
use rustg_cuda::prelude::*;
#[kernel]
pub fn shared_memory_kernel(input: &[f32], output: &mut [f32]) {
// Allocate shared memory
let mut shared: [f32; 512] = [0.0; 512];
let tid = thread_idx_x();
// Load data into shared memory
if tid < input.len() {
shared[tid] = input[tid];
}
sync_threads();
// Process data in shared memory
if tid < output.len() {
output[tid] = shared[tid] * 2.0;
}
}
"#;
let options = CompileOptions {
target: Target::SM120,
optimize: true,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
let result = compiler.compile_kernel("shared_mem_test", shared_memory_kernel);
assert!(result.is_ok(), "Shared memory kernel should compile");
let ptx = result.unwrap();
let ptx_str = String::from_utf8_lossy(&ptx);
// Check for shared memory declarations in PTX
assert!(
ptx_str.contains(".shared") || ptx_str.len() > 300,
"PTX should contain shared memory optimizations"
);
}
#[test]
fn test_register_allocation_limits() {
let register_heavy_kernel = r#"
#![no_std]
use rustg_cuda::prelude::*;
#[kernel]
pub fn register_heavy_kernel(data: &mut [f32]) {
let idx = thread_idx_x();
// Use many variables to stress register allocation
let mut vars = [0.0f32; 32];
for i in 0..32 {
vars[i] = (idx as f32) * (i as f32);
}
if idx < data.len() {
data[idx] = vars.iter().sum();
}
}
"#;
let options = CompileOptions {
target: Target::SM120,
max_registers: Some(64),
optimize: true,
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
let result = compiler.compile_kernel("register_test", register_heavy_kernel);
assert!(
result.is_ok(),
"Register-heavy kernel should compile with limits"
);
let ptx = result.unwrap();
let ptx_str = String::from_utf8_lossy(&ptx);
// Check that register limit is respected
assert!(
ptx_str.contains(".maxnreg 64") || ptx_str.len() > 200,
"PTX should respect register allocation limits"
);
}
#[test]
fn test_incremental_compilation() {
let cache_dir = std::env::temp_dir().join("rtx_incremental_test");
let _ = fs::remove_dir_all(&cache_dir);
let options = CompileOptions {
cache_dir: cache_dir.clone(),
..Default::default()
};
let mut compiler = RtxCompiler::new(options);
// Compile original version
let original_source = SIMPLE_KERNEL_SOURCE;
let result1 = compiler.compile_kernel("incremental_test", original_source);
assert!(result1.is_ok(), "Original compilation should succeed");
// Modify source slightly
let modified_source = original_source.replace(
"c[idx] = a[idx] + b[idx];",
"c[idx] = a[idx] + b[idx] + 0.1;",
);
// Second compilation should detect change and recompile
let result2 = compiler.compile_kernel("incremental_test", &modified_source);
assert!(result2.is_ok(), "Modified compilation should succeed");
// Results should be different
assert_ne!(
result1.unwrap(),
result2.unwrap(),
"Modified kernel should produce different output"
);
let _ = fs::remove_dir_all(&cache_dir);
}
#[test]
fn test_rustg_binary_availability() {
// Test that rustg binary is available or graceful fallback occurs
let mut compiler = RtxCompiler::default();
// This should either:
// 1. Successfully compile with rustg if available
// 2. Provide a meaningful error about rustg not being available
// 3. Fall back to a mock implementation for development
let result = compiler.compile_kernel("availability_test", SIMPLE_KERNEL_SOURCE);
if result.is_err() {
match result.unwrap_err() {
CompilerError::RustgCompilationFailed(msg) => {
// Should provide helpful error about rustg availability
assert!(
msg.contains("rustg")
|| msg.contains("not found")
|| msg.contains("binary"),
"Error should mention rustg availability: {}",
msg
);
}
_ => {
// Other errors are acceptable for this test
}
}
} else {
// Success is also acceptable
let ptx = result.unwrap();
assert!(
!ptx.is_empty(),
"Successful compilation should produce output"
);
}
}
}