use crate::error::{InferenceError, InferenceResult};
use half::{bf16, f16};
use scirs2_core::ndarray::{Array1, Array2};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)]
pub enum PrecisionMode {
#[default]
FP32,
FP16,
BF16,
Mixed {
compute: ComputePrecision,
accumulate_fp32: bool,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum ComputePrecision {
FP16,
BF16,
}
impl PrecisionMode {
pub fn is_reduced_precision(&self) -> bool {
!matches!(self, PrecisionMode::FP32)
}
pub fn memory_reduction_factor(&self) -> f32 {
match self {
PrecisionMode::FP32 => 1.0,
PrecisionMode::FP16 | PrecisionMode::BF16 => 0.5,
PrecisionMode::Mixed { .. } => 0.75, }
}
pub fn name(&self) -> &str {
match self {
PrecisionMode::FP32 => "FP32",
PrecisionMode::FP16 => "FP16",
PrecisionMode::BF16 => "BF16",
PrecisionMode::Mixed { compute, .. } => match compute {
ComputePrecision::FP16 => "Mixed-FP16",
ComputePrecision::BF16 => "Mixed-BF16",
},
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct PrecisionConfig {
pub mode: PrecisionMode,
pub loss_scale: f32,
pub dynamic_loss_scale: bool,
pub grad_clip_threshold: Option<f32>,
}
impl Default for PrecisionConfig {
fn default() -> Self {
Self {
mode: PrecisionMode::FP32,
loss_scale: 1.0,
dynamic_loss_scale: false,
grad_clip_threshold: None,
}
}
}
impl PrecisionConfig {
pub fn new() -> Self {
Self::default()
}
pub fn mode(mut self, mode: PrecisionMode) -> Self {
self.mode = mode;
self
}
pub fn fp16(mut self) -> Self {
self.mode = PrecisionMode::FP16;
self
}
pub fn bf16(mut self) -> Self {
self.mode = PrecisionMode::BF16;
self
}
pub fn mixed_fp16(mut self, accumulate_fp32: bool) -> Self {
self.mode = PrecisionMode::Mixed {
compute: ComputePrecision::FP16,
accumulate_fp32,
};
self
}
pub fn mixed_bf16(mut self, accumulate_fp32: bool) -> Self {
self.mode = PrecisionMode::Mixed {
compute: ComputePrecision::BF16,
accumulate_fp32,
};
self
}
pub fn loss_scale(mut self, scale: f32) -> Self {
self.loss_scale = scale;
self
}
pub fn dynamic_loss_scale(mut self, enabled: bool) -> Self {
self.dynamic_loss_scale = enabled;
self
}
pub fn grad_clip_threshold(mut self, threshold: f32) -> Self {
self.grad_clip_threshold = Some(threshold);
self
}
}
pub struct PrecisionConverter {
config: PrecisionConfig,
}
impl PrecisionConverter {
pub fn new(config: PrecisionConfig) -> Self {
Self { config }
}
pub fn convert_and_compute_1d(
&self,
data: &Array1<f32>,
op: impl Fn(&Array1<f32>) -> Array1<f32>,
) -> InferenceResult<Array1<f32>> {
match self.config.mode {
PrecisionMode::FP32 => Ok(op(data)),
PrecisionMode::FP16 => {
let fp16_data = self.to_fp16_1d(data);
let fp16_result = op(&self.from_fp16_1d(&fp16_data));
Ok(fp16_result)
}
PrecisionMode::BF16 => {
let bf16_data = self.to_bf16_1d(data);
let bf16_result = op(&self.from_bf16_1d(&bf16_data));
Ok(bf16_result)
}
PrecisionMode::Mixed {
compute,
accumulate_fp32,
} => {
if accumulate_fp32 {
let reduced = match compute {
ComputePrecision::FP16 => {
let fp16_data = self.to_fp16_1d(data);
self.from_fp16_1d(&fp16_data)
}
ComputePrecision::BF16 => {
let bf16_data = self.to_bf16_1d(data);
self.from_bf16_1d(&bf16_data)
}
};
Ok(op(&reduced))
} else {
match compute {
ComputePrecision::FP16 => {
let fp16_data = self.to_fp16_1d(data);
Ok(op(&self.from_fp16_1d(&fp16_data)))
}
ComputePrecision::BF16 => {
let bf16_data = self.to_bf16_1d(data);
Ok(op(&self.from_bf16_1d(&bf16_data)))
}
}
}
}
}
}
pub fn to_fp16_1d(&self, data: &Array1<f32>) -> Vec<f16> {
data.iter().map(|&x| f16::from_f32(x)).collect()
}
pub fn from_fp16_1d(&self, data: &[f16]) -> Array1<f32> {
Array1::from_vec(data.iter().map(|&x| x.to_f32()).collect())
}
pub fn to_bf16_1d(&self, data: &Array1<f32>) -> Vec<bf16> {
data.iter().map(|&x| bf16::from_f32(x)).collect()
}
pub fn from_bf16_1d(&self, data: &[bf16]) -> Array1<f32> {
Array1::from_vec(data.iter().map(|&x| x.to_f32()).collect())
}
pub fn to_fp16_2d(&self, data: &Array2<f32>) -> Vec<f16> {
data.iter().map(|&x| f16::from_f32(x)).collect()
}
pub fn from_fp16_2d(
&self,
data: &[f16],
shape: (usize, usize),
) -> InferenceResult<Array2<f32>> {
let vec: Vec<f32> = data.iter().map(|&x| x.to_f32()).collect();
Array2::from_shape_vec(shape, vec).map_err(|e| {
InferenceError::ForwardError(format!("Shape error in FP16 conversion: {}", e))
})
}
pub fn to_bf16_2d(&self, data: &Array2<f32>) -> Vec<bf16> {
data.iter().map(|&x| bf16::from_f32(x)).collect()
}
pub fn from_bf16_2d(
&self,
data: &[bf16],
shape: (usize, usize),
) -> InferenceResult<Array2<f32>> {
let vec: Vec<f32> = data.iter().map(|&x| x.to_f32()).collect();
Array2::from_shape_vec(shape, vec).map_err(|e| {
InferenceError::ForwardError(format!("Shape error in BF16 conversion: {}", e))
})
}
pub fn config(&self) -> &PrecisionConfig {
&self.config
}
}
#[derive(Debug, Clone, Default)]
pub struct PrecisionStats {
pub num_conversions: usize,
pub memory_saved: usize,
pub avg_error: f64,
pub max_error: f64,
}
impl PrecisionStats {
pub fn new() -> Self {
Self::default()
}
pub fn record_conversion(&mut self, original_size: usize, precision_mode: &PrecisionMode) {
self.num_conversions += 1;
let saved =
(original_size as f32 * (1.0 - precision_mode.memory_reduction_factor())) as usize;
self.memory_saved += saved;
}
pub fn record_error(&mut self, error: f64) {
let n = self.num_conversions as f64;
self.avg_error = (self.avg_error * (n - 1.0) + error) / n;
self.max_error = self.max_error.max(error);
}
pub fn memory_saved_mb(&self) -> f64 {
self.memory_saved as f64 / (1024.0 * 1024.0)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_precision_mode_creation() {
let mode = PrecisionMode::FP32;
assert_eq!(mode.name(), "FP32");
assert!(!mode.is_reduced_precision());
}
#[test]
fn test_precision_mode_fp16() {
let mode = PrecisionMode::FP16;
assert_eq!(mode.name(), "FP16");
assert!(mode.is_reduced_precision());
assert_eq!(mode.memory_reduction_factor(), 0.5);
}
#[test]
fn test_precision_mode_bf16() {
let mode = PrecisionMode::BF16;
assert_eq!(mode.name(), "BF16");
assert!(mode.is_reduced_precision());
assert_eq!(mode.memory_reduction_factor(), 0.5);
}
#[test]
fn test_precision_mode_mixed() {
let mode = PrecisionMode::Mixed {
compute: ComputePrecision::FP16,
accumulate_fp32: true,
};
assert_eq!(mode.name(), "Mixed-FP16");
assert!(mode.is_reduced_precision());
}
#[test]
fn test_precision_config_builder() {
let config = PrecisionConfig::new()
.fp16()
.loss_scale(128.0)
.dynamic_loss_scale(true);
assert_eq!(config.mode, PrecisionMode::FP16);
assert_eq!(config.loss_scale, 128.0);
assert!(config.dynamic_loss_scale);
}
#[test]
fn test_fp16_conversion_1d() {
let config = PrecisionConfig::new().fp16();
let converter = PrecisionConverter::new(config);
let data = Array1::from_vec(vec![1.0, 2.5, -3.75, 0.0]);
let fp16_data = converter.to_fp16_1d(&data);
let restored = converter.from_fp16_1d(&fp16_data);
for (orig, rest) in data.iter().zip(restored.iter()) {
assert!((orig - rest).abs() < 0.001);
}
}
#[test]
fn test_bf16_conversion_1d() {
let config = PrecisionConfig::new().bf16();
let converter = PrecisionConverter::new(config);
let data = Array1::from_vec(vec![1.0, 2.5, -3.75, 0.0]);
let bf16_data = converter.to_bf16_1d(&data);
let restored = converter.from_bf16_1d(&bf16_data);
for (orig, rest) in data.iter().zip(restored.iter()) {
assert!((orig - rest).abs() < 0.01);
}
}
#[test]
fn test_convert_and_compute() {
let config = PrecisionConfig::new().fp16();
let converter = PrecisionConverter::new(config);
let data = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let result = converter
.convert_and_compute_1d(&data, |x| x.mapv(|v| v * 2.0))
.unwrap();
for (i, &val) in result.iter().enumerate() {
let expected = data[i] * 2.0;
assert!((val - expected).abs() < 0.01);
}
}
#[test]
fn test_precision_stats() {
let mut stats = PrecisionStats::new();
assert_eq!(stats.num_conversions, 0);
assert_eq!(stats.memory_saved, 0);
let mode = PrecisionMode::FP16;
stats.record_conversion(1000, &mode);
assert_eq!(stats.num_conversions, 1);
assert_eq!(stats.memory_saved, 500);
stats.record_error(0.001);
assert!(stats.avg_error > 0.0);
assert!(stats.max_error > 0.0);
}
#[test]
fn test_mixed_precision_compute() {
let config = PrecisionConfig::new().mixed_fp16(true);
let converter = PrecisionConverter::new(config);
let data = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let result = converter
.convert_and_compute_1d(&data, |x| x.mapv(|v| v * 2.0))
.unwrap();
for (i, &val) in result.iter().enumerate() {
let expected = data[i] * 2.0;
assert!((val - expected).abs() < 0.001);
}
}
#[test]
fn test_fp16_2d_conversion() {
let config = PrecisionConfig::new().fp16();
let converter = PrecisionConverter::new(config);
let data = Array2::from_shape_vec((2, 3), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
let fp16_data = converter.to_fp16_2d(&data);
let restored = converter.from_fp16_2d(&fp16_data, (2, 3)).unwrap();
assert_eq!(restored.shape(), &[2, 3]);
for (orig, rest) in data.iter().zip(restored.iter()) {
assert!((orig - rest).abs() < 0.001);
}
}
#[test]
fn test_bf16_2d_conversion() {
let config = PrecisionConfig::new().bf16();
let converter = PrecisionConverter::new(config);
let data = Array2::from_shape_vec((2, 2), vec![1.0, 2.0, 3.0, 4.0]).unwrap();
let bf16_data = converter.to_bf16_2d(&data);
let restored = converter.from_bf16_2d(&bf16_data, (2, 2)).unwrap();
assert_eq!(restored.shape(), &[2, 2]);
for (orig, rest) in data.iter().zip(restored.iter()) {
assert!((orig - rest).abs() < 0.01);
}
}
}