use crate::bridge;
use alloc::boxed::Box;
use alloc::format;
use alloc::string::String;
use alloc::string::ToString;
use alloc::vec;
use burn_pack::Tensor as PackTensor;
use burn_core::tensor::shape;
use burn_core::tensor::{DType, TensorData};
use hashbrown::HashSet;
mod module_names {
pub const LINEAR: &str = "Struct:Linear";
pub const BATCH_NORM: &str = "Struct:BatchNorm";
pub const LAYER_NORM: &str = "Struct:LayerNorm";
pub const GROUP_NORM: &str = "Struct:GroupNorm";
pub const EMBEDDING: &str = "Struct:Embedding";
pub const CONV1D: &str = "Struct:Conv1d";
pub const CONV2D: &str = "Struct:Conv2d";
pub const CONV3D: &str = "Struct:Conv3d";
pub const CONV_TRANSPOSE1D: &str = "Struct:ConvTranspose1d";
pub const CONV_TRANSPOSE2D: &str = "Struct:ConvTranspose2d";
pub const CONV_TRANSPOSE3D: &str = "Struct:ConvTranspose3d";
pub const DEFORM_CONV2D: &str = "Struct:DeformConv2d";
pub const INSTANCE_NORM: &str = "Struct:InstanceNorm";
pub const RMS_NORM: &str = "Struct:RmsNorm";
pub const PRELU: &str = "Struct:PRelu";
}
#[derive(Debug, Clone, Copy)]
pub struct ModuleContext<'a> {
container_stack: &'a [String],
}
impl<'a> ModuleContext<'a> {
pub fn new(container_stack: &'a [String]) -> Self {
Self { container_stack }
}
pub fn module_type(&self) -> Option<&'a str> {
self.container_stack
.iter()
.rev()
.find(|ct| ct.starts_with("Struct:") || ct.starts_with("Enum:"))
.map(|s| s.as_str())
}
}
pub trait ModuleAdapter: Send + Sync {
fn adapt(&self, tensor: PackTensor, ctx: ModuleContext<'_>) -> PackTensor;
fn get_alternative_param_name(&self, _param_name: &str, _module_type: &str) -> Option<String> {
None
}
fn clone_box(&self) -> Box<dyn ModuleAdapter>;
fn chain<A>(self, next: A) -> ChainAdapter
where
Self: Sized + 'static,
A: ModuleAdapter + 'static,
{
ChainAdapter::new(self, next)
}
}
impl Clone for Box<dyn ModuleAdapter> {
fn clone(&self) -> Self {
self.clone_box()
}
}
#[derive(Clone)]
pub struct ChainAdapter {
first: Box<dyn ModuleAdapter>,
second: Box<dyn ModuleAdapter>,
}
impl ChainAdapter {
pub fn new<A, B>(first: A, second: B) -> Self
where
A: ModuleAdapter + 'static,
B: ModuleAdapter + 'static,
{
Self {
first: Box::new(first),
second: Box::new(second),
}
}
}
impl ModuleAdapter for ChainAdapter {
fn adapt(&self, tensor: PackTensor, ctx: ModuleContext<'_>) -> PackTensor {
self.second.adapt(self.first.adapt(tensor, ctx), ctx)
}
fn get_alternative_param_name(&self, param_name: &str, container_type: &str) -> Option<String> {
if let Some(name) = self
.first
.get_alternative_param_name(param_name, container_type)
{
self.second
.get_alternative_param_name(&name, container_type)
.or(Some(name))
} else {
self.second
.get_alternative_param_name(param_name, container_type)
}
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Debug, Clone, Default)]
pub struct IdentityAdapter;
impl ModuleAdapter for IdentityAdapter {
fn adapt(&self, tensor: PackTensor, _ctx: ModuleContext<'_>) -> PackTensor {
tensor
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
fn cast(tensor: PackTensor, target: DType) -> PackTensor {
let (name, shape) = (tensor.name.clone(), tensor.shape.clone());
bridge::map_data(tensor, name, target, shape, move |data| {
data.convert_dtype(target)
})
}
fn param_name(tensor: &PackTensor) -> &str {
tensor.name.rsplit('.').next().unwrap_or("")
}
fn rename_param(mut tensor: PackTensor, new_name: &str) -> PackTensor {
let keep = tensor.name.rfind('.').map(|i| i + 1).unwrap_or(0);
tensor.name.truncate(keep);
tensor.name.push_str(new_name);
tensor
}
fn default_half_precision_modules() -> HashSet<String> {
let modules = [
module_names::LINEAR,
module_names::EMBEDDING,
module_names::CONV1D,
module_names::CONV2D,
module_names::CONV3D,
module_names::CONV_TRANSPOSE1D,
module_names::CONV_TRANSPOSE2D,
module_names::CONV_TRANSPOSE3D,
module_names::DEFORM_CONV2D,
module_names::LAYER_NORM,
module_names::GROUP_NORM,
module_names::INSTANCE_NORM,
module_names::RMS_NORM,
module_names::PRELU,
];
modules.iter().map(|s| s.to_string()).collect()
}
#[derive(Debug, Clone)]
pub struct HalfPrecisionAdapter {
modules: HashSet<String>,
}
impl HalfPrecisionAdapter {
pub fn new() -> Self {
Self {
modules: default_half_precision_modules(),
}
}
pub fn with_module(mut self, module_type: impl Into<String>) -> Self {
let name = module_type.into();
if name.contains(':') {
self.modules.insert(name);
} else {
self.modules.insert(format!("Struct:{}", name));
}
self
}
pub fn without_module(mut self, module_type: impl Into<String>) -> Self {
let name = module_type.into();
let key = if name.contains(':') {
name
} else {
format!("Struct:{}", name)
};
assert!(
self.modules.contains(&key),
"without_module called with '{}' which is not in the module set",
key
);
self.modules.remove(&key);
self
}
fn should_convert(&self, ctx: ModuleContext<'_>) -> bool {
ctx.module_type()
.is_some_and(|mt| self.modules.contains(mt))
}
}
impl Default for HalfPrecisionAdapter {
fn default() -> Self {
Self::new()
}
}
impl ModuleAdapter for HalfPrecisionAdapter {
fn adapt(&self, tensor: PackTensor, ctx: ModuleContext<'_>) -> PackTensor {
let target_dtype = match tensor.dtype {
DType::F32 => DType::F16,
DType::F16 => DType::F32,
_ => return tensor,
};
if !self.should_convert(ctx) {
return tensor;
}
cast(tensor, target_dtype)
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Debug, Clone)]
pub struct FloatCastAdapter {
target: DType,
}
impl FloatCastAdapter {
pub fn to(target: DType) -> Self {
assert!(
target.is_float(),
"FloatCastAdapter target must be a float dtype, got {target:?}"
);
Self { target }
}
}
impl ModuleAdapter for FloatCastAdapter {
fn adapt(&self, tensor: PackTensor, _ctx: ModuleContext<'_>) -> PackTensor {
if !tensor.dtype.is_float() || tensor.dtype == self.target {
return tensor;
}
cast(tensor, self.target)
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Debug, Clone, Default)]
pub struct PyTorchToBurnAdapter;
impl ModuleAdapter for PyTorchToBurnAdapter {
fn adapt(&self, tensor: PackTensor, ctx: ModuleContext<'_>) -> PackTensor {
adapt_pytorch_tensor(tensor, ctx, PyTorchConversionDirection::PyTorchToBurn)
}
fn get_alternative_param_name(&self, param_name: &str, container_type: &str) -> Option<String> {
if is_normalization_layer(container_type) {
burn_norm_param_to_pytorch(param_name).map(|s| s.to_string())
} else {
None
}
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Debug, Clone, Default)]
pub struct BurnToPyTorchAdapter;
impl ModuleAdapter for BurnToPyTorchAdapter {
fn adapt(&self, tensor: PackTensor, ctx: ModuleContext<'_>) -> PackTensor {
adapt_pytorch_tensor(tensor, ctx, PyTorchConversionDirection::BurnToPyTorch)
}
fn get_alternative_param_name(&self, param_name: &str, container_type: &str) -> Option<String> {
if is_normalization_layer(container_type) {
pytorch_norm_param_to_burn(param_name).map(|s| s.to_string())
} else {
None
}
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Debug, Clone, Copy)]
enum PyTorchConversionDirection {
PyTorchToBurn,
BurnToPyTorch,
}
fn is_normalization_layer(container_type: &str) -> bool {
matches!(
container_type,
module_names::BATCH_NORM
| module_names::LAYER_NORM
| module_names::GROUP_NORM
| module_names::RMS_NORM
)
}
fn pytorch_norm_param_to_burn(param_name: &str) -> Option<&'static str> {
match param_name {
"weight" => Some("gamma"),
"bias" => Some("beta"),
_ => None,
}
}
fn burn_norm_param_to_pytorch(param_name: &str) -> Option<&'static str> {
match param_name {
"gamma" => Some("weight"),
"beta" => Some("bias"),
_ => None,
}
}
fn adapt_pytorch_tensor(
tensor: PackTensor,
ctx: ModuleContext<'_>,
direction: PyTorchConversionDirection,
) -> PackTensor {
let Some(module_type) = ctx.module_type() else {
return tensor; };
let param = param_name(&tensor);
if module_type == module_names::LINEAR && param == "weight" && tensor.shape.len() == 2 {
return transpose_2d_tensor(tensor);
}
if is_normalization_layer(module_type) {
let renamed = match direction {
PyTorchConversionDirection::PyTorchToBurn => pytorch_norm_param_to_burn(param),
PyTorchConversionDirection::BurnToPyTorch => burn_norm_param_to_pytorch(param),
};
if let Some(new_name) = renamed {
return rename_param(tensor, new_name);
}
}
tensor
}
fn transpose_2d_tensor(tensor: PackTensor) -> PackTensor {
if tensor.shape.len() != 2 {
return tensor;
}
if matches!(tensor.dtype, DType::QFloat(_)) {
return tensor;
}
let (name, dtype) = (tensor.name.clone(), tensor.dtype);
let transposed_shape = shape![tensor.shape[1], tensor.shape[0]];
bridge::map_data(tensor, name, dtype, transposed_shape, transpose_tensor_data)
}
fn transpose_tensor_data(data: TensorData) -> TensorData {
let shape = &data.shape;
let rows = shape[0];
let cols = shape[1];
let transposed_shape = vec![cols, rows];
let bytes = data.as_bytes();
let element_size = data.dtype.size();
let mut transposed_bytes = vec![0u8; bytes.len()];
for i in 0..rows {
for j in 0..cols {
let src_idx = (i * cols + j) * element_size;
let dst_idx = (j * rows + i) * element_size;
transposed_bytes[dst_idx..dst_idx + element_size]
.copy_from_slice(&bytes[src_idx..src_idx + element_size]);
}
}
TensorData::from_bytes_vec(transposed_bytes, transposed_shape, data.dtype)
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::string::ToString;
use alloc::sync::Arc;
use alloc::vec::Vec;
use burn_core::tensor::{Bytes, DType, Shape, TensorData};
use core::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn test_module_names_match_burn_nn() {
#[allow(unused_imports)]
use burn_nn::{
BatchNorm, Embedding, GroupNorm, InstanceNorm, LayerNorm, Linear, PRelu, RmsNorm,
conv::{
Conv1d, Conv2d, Conv3d, ConvTranspose1d, ConvTranspose2d, ConvTranspose3d,
DeformConv2d,
},
};
assert_eq!(module_names::LINEAR, "Struct:Linear");
assert_eq!(module_names::BATCH_NORM, "Struct:BatchNorm");
assert_eq!(module_names::LAYER_NORM, "Struct:LayerNorm");
assert_eq!(module_names::GROUP_NORM, "Struct:GroupNorm");
assert_eq!(module_names::EMBEDDING, "Struct:Embedding");
assert_eq!(module_names::CONV1D, "Struct:Conv1d");
assert_eq!(module_names::CONV2D, "Struct:Conv2d");
assert_eq!(module_names::CONV3D, "Struct:Conv3d");
assert_eq!(module_names::CONV_TRANSPOSE1D, "Struct:ConvTranspose1d");
assert_eq!(module_names::CONV_TRANSPOSE2D, "Struct:ConvTranspose2d");
assert_eq!(module_names::CONV_TRANSPOSE3D, "Struct:ConvTranspose3d");
assert_eq!(module_names::DEFORM_CONV2D, "Struct:DeformConv2d");
assert_eq!(module_names::INSTANCE_NORM, "Struct:InstanceNorm");
assert_eq!(module_names::RMS_NORM, "Struct:RmsNorm");
assert_eq!(module_names::PRELU, "Struct:PRelu");
}
fn tensor_with(name: &str, data: TensorData) -> PackTensor {
bridge::from_data(data, name.to_string(), None)
}
fn tensor(name: &str, shape: Shape) -> PackTensor {
let values = vec![1.0f32; shape.iter().product()];
tensor_with(name, TensorData::new(values, shape))
}
fn containers(container_type: &str) -> Vec<String> {
vec![container_type.to_string()]
}
fn adapt_in(
adapter: &dyn ModuleAdapter,
name: &str,
shape: Shape,
container_type: &str,
) -> PackTensor {
let containers = containers(container_type);
adapter.adapt(tensor(name, shape), ModuleContext::new(&containers))
}
fn adapt_without_module(adapter: &dyn ModuleAdapter, tensor: PackTensor) -> PackTensor {
adapter.adapt(tensor, ModuleContext::new(&[]))
}
#[test]
fn test_pytorch_to_burn_linear_weight() {
let adapter = PyTorchToBurnAdapter;
let adapted = adapt_in(&adapter, "fc.weight", shape![10, 5], module_names::LINEAR);
assert_eq!(adapted.shape, shape![5, 10]);
let adapted = adapt_in(&adapter, "fc.bias", shape![10], module_names::LINEAR);
assert_eq!(adapted.shape, shape![10]);
}
#[test]
fn test_pytorch_to_burn_norm_params() {
let adapter = PyTorchToBurnAdapter;
let adapted = adapt_in(
&adapter,
"norm.weight",
shape![10],
module_names::BATCH_NORM,
);
assert_eq!(adapted.name, "norm.gamma");
let adapted = adapt_in(&adapter, "norm.bias", shape![10], module_names::BATCH_NORM);
assert_eq!(adapted.name, "norm.beta");
}
#[test]
fn test_burn_to_pytorch_linear_weight() {
let adapter = BurnToPyTorchAdapter;
let adapted = adapt_in(&adapter, "fc.weight", shape![5, 10], module_names::LINEAR);
assert_eq!(adapted.shape, shape![10, 5]);
}
#[test]
fn test_burn_to_pytorch_norm_params() {
let adapter = BurnToPyTorchAdapter;
let adapted = adapt_in(&adapter, "norm.gamma", shape![10], module_names::BATCH_NORM);
assert_eq!(adapted.name, "norm.weight");
let adapted = adapt_in(&adapter, "norm.beta", shape![10], module_names::BATCH_NORM);
assert_eq!(adapted.name, "norm.bias");
}
#[test]
fn rename_keeps_the_enclosing_path() {
let adapter = PyTorchToBurnAdapter;
let adapted = adapt_in(
&adapter,
"encoder.layers.0.norm.weight",
shape![10],
module_names::LAYER_NORM,
);
assert_eq!(adapted.name, "encoder.layers.0.norm.gamma");
let adapted = adapt_in(&adapter, "weight", shape![10], module_names::LAYER_NORM);
assert_eq!(adapted.name, "gamma");
}
#[test]
fn test_transpose_different_dtypes() {
let f32_data = TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], [2, 3]);
let transposed = transpose_tensor_data(f32_data);
assert_eq!(transposed.shape, shape![3, 2]);
let values = transposed.try_to_vec::<f32>().unwrap();
assert_eq!(values, vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]);
let i32_data = TensorData::new(vec![1i32, 2, 3, 4, 5, 6], [2, 3]);
let transposed = transpose_tensor_data(i32_data);
assert_eq!(transposed.shape, shape![3, 2]);
let values = transposed.try_to_vec::<i32>().unwrap();
assert_eq!(values, vec![1, 4, 2, 5, 3, 6]);
let f64_data = TensorData::new(vec![1.0f64, 2.0, 3.0, 4.0], [2, 2]);
let transposed = transpose_tensor_data(f64_data);
assert_eq!(transposed.shape, shape![2, 2]);
let values = transposed.try_to_vec::<f64>().unwrap();
assert_eq!(values, vec![1.0, 3.0, 2.0, 4.0]);
}
#[test]
fn transpose_moves_the_data_not_just_the_shape() {
let adapter = PyTorchToBurnAdapter;
let data = TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], [2, 3]);
let containers = containers(module_names::LINEAR);
let adapted = adapter.adapt(
tensor_with("fc.weight", data),
ModuleContext::new(&containers),
);
assert_eq!(adapted.shape, shape![3, 2]);
let values = bridge::to_data(&adapted)
.unwrap()
.try_into_vec::<f32>()
.unwrap();
assert_eq!(values, vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]);
}
#[test]
fn test_no_container_info() {
let adapter = PyTorchToBurnAdapter;
let adapted = adapt_without_module(&adapter, tensor("fc.weight", shape![10, 5]));
assert_eq!(adapted.shape, shape![10, 5]);
let adapted = adapt_without_module(&adapter, tensor("other.weight", shape![10, 5]));
assert_eq!(adapted.shape, shape![10, 5]); }
#[test]
fn module_type_looks_past_collection_wrappers() {
let adapter = PyTorchToBurnAdapter;
let containers = vec![
"Struct:Model".to_string(),
"Vec".to_string(),
module_names::LINEAR.to_string(),
];
let adapted = adapter.adapt(
tensor("layers.0.weight", shape![10, 5]),
ModuleContext::new(&containers),
);
assert_eq!(adapted.shape, shape![5, 10]);
}
#[derive(Clone)]
struct RenameParamAdapter {
from: &'static str,
to: &'static str,
called: Arc<AtomicUsize>,
}
impl ModuleAdapter for RenameParamAdapter {
fn adapt(&self, tensor: PackTensor, _ctx: ModuleContext<'_>) -> PackTensor {
self.called.fetch_add(1, Ordering::Relaxed);
if param_name(&tensor) != self.from {
return tensor;
}
rename_param(tensor, self.to)
}
fn get_alternative_param_name(
&self,
_param_name: &str,
_container_type: &str,
) -> Option<String> {
None
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Clone)]
struct AltNameAdapter {
from: &'static str,
to: &'static str,
called: Arc<AtomicUsize>,
}
impl ModuleAdapter for AltNameAdapter {
fn adapt(&self, tensor: PackTensor, _ctx: ModuleContext<'_>) -> PackTensor {
tensor
}
fn get_alternative_param_name(
&self,
param_name: &str,
_container_type: &str,
) -> Option<String> {
self.called.fetch_add(1, Ordering::Relaxed);
if param_name == self.from {
Some(self.to.to_string())
} else {
None
}
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[test]
fn test_chain_adapter_pipes_adapt() {
let called1 = Arc::new(AtomicUsize::new(0));
let called2 = Arc::new(AtomicUsize::new(0));
let a = RenameParamAdapter {
from: "weight",
to: "a",
called: called1.clone(),
};
let b = RenameParamAdapter {
from: "a",
to: "b",
called: called2.clone(),
};
let chain = a.chain(b);
let adapted = adapt_in(&chain, "fc.weight", shape![2, 2], module_names::LINEAR);
assert_eq!(adapted.name, "fc.b");
assert_eq!(called1.load(Ordering::Relaxed), 1);
assert_eq!(called2.load(Ordering::Relaxed), 1);
}
#[test]
fn test_chain_adapter_alternative_name_pipes_and_fallbacks() {
let called1 = Arc::new(AtomicUsize::new(0));
let called2 = Arc::new(AtomicUsize::new(0));
let a = AltNameAdapter {
from: "gamma",
to: "weight",
called: called1.clone(),
};
let b = AltNameAdapter {
from: "weight",
to: "scale",
called: called2.clone(),
};
let chain = a.chain(b);
let alt = chain.get_alternative_param_name("gamma", module_names::LAYER_NORM);
assert_eq!(alt.as_deref(), Some("scale"));
assert_eq!(called1.load(Ordering::Relaxed), 1);
assert_eq!(called2.load(Ordering::Relaxed), 1);
let called1 = Arc::new(AtomicUsize::new(0));
let called2 = Arc::new(AtomicUsize::new(0));
let a = AltNameAdapter {
from: "gamma",
to: "weight",
called: called1.clone(),
};
let b = AltNameAdapter {
from: "something-else",
to: "unused",
called: called2.clone(),
};
let chain = a.chain(b);
let alt = chain.get_alternative_param_name("gamma", module_names::LAYER_NORM);
assert_eq!(alt.as_deref(), Some("weight"));
assert_eq!(called1.load(Ordering::Relaxed), 1);
assert_eq!(called2.load(Ordering::Relaxed), 1);
let called1 = Arc::new(AtomicUsize::new(0));
let called2 = Arc::new(AtomicUsize::new(0));
let a = AltNameAdapter {
from: "something-else",
to: "unused",
called: called1.clone(),
};
let b = AltNameAdapter {
from: "gamma",
to: "weight",
called: called2.clone(),
};
let chain = a.chain(b);
let alt = chain.get_alternative_param_name("gamma", module_names::LAYER_NORM);
assert_eq!(alt.as_deref(), Some("weight"));
assert_eq!(called1.load(Ordering::Relaxed), 1);
assert_eq!(called2.load(Ordering::Relaxed), 1);
let boxed = chain.clone_box();
let alt = boxed.get_alternative_param_name("gamma", module_names::LAYER_NORM);
assert_eq!(alt.as_deref(), Some("weight"));
}
#[test]
fn test_half_precision_f32_to_f16() {
let adapter = HalfPrecisionAdapter::new();
let adapted = adapt_in(&adapter, "fc.weight", shape![2, 3], module_names::LINEAR);
assert_eq!(adapted.dtype, DType::F16);
assert_eq!(adapted.shape, shape![2, 3]);
let data = bridge::to_data(&adapted).unwrap();
assert_eq!(data.dtype, DType::F16);
}
#[test]
fn test_half_precision_updates_byte_len() {
let adapter = HalfPrecisionAdapter::new();
let adapted = adapt_in(&adapter, "fc.weight", shape![2, 3], module_names::LINEAR);
assert_eq!(adapted.byte_len(), 6 * 2);
assert_eq!(bridge::to_data(&adapted).unwrap().bytes.len(), 6 * 2);
}
#[test]
fn test_half_precision_f16_to_f32() {
let adapter = HalfPrecisionAdapter::new();
let data = TensorData::new(vec![1.0f32; 6], shape![2, 3]).convert_dtype(DType::F16);
let containers = containers(module_names::LINEAR);
let adapted = adapter.adapt(
tensor_with("fc.weight", data),
ModuleContext::new(&containers),
);
assert_eq!(adapted.dtype, DType::F32);
}
#[test]
fn test_half_precision_skips_batch_norm() {
let adapter = HalfPrecisionAdapter::new();
let adapted = adapt_in(
&adapter,
"norm.weight",
shape![10],
module_names::BATCH_NORM,
);
assert_eq!(adapted.dtype, DType::F32); }
#[test]
fn test_half_precision_converts_default_modules() {
let adapter = HalfPrecisionAdapter::new();
for (name, shape, module) in [
("fc.weight", shape![2, 3], module_names::LINEAR),
("emb.weight", shape![100, 64], module_names::EMBEDDING),
("conv.weight", shape![3, 3, 3, 3], module_names::CONV2D),
("norm.gamma", shape![10], module_names::LAYER_NORM),
("gn.gamma", shape![10], module_names::GROUP_NORM),
("rms.weight", shape![10], module_names::RMS_NORM),
] {
assert_eq!(
adapt_in(&adapter, name, shape, module).dtype,
DType::F16,
"{module} should be converted"
);
}
}
#[test]
fn test_half_precision_without_module() {
let adapter = HalfPrecisionAdapter::new().without_module("LayerNorm");
let adapted = adapt_in(&adapter, "norm.gamma", shape![10], module_names::LAYER_NORM);
assert_eq!(adapted.dtype, DType::F32);
let adapted = adapt_in(&adapter, "fc.weight", shape![2, 3], module_names::LINEAR);
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
fn test_half_precision_with_module() {
let adapter = HalfPrecisionAdapter::new().with_module("CustomLayer");
let adapted = adapt_in(&adapter, "custom.weight", shape![5], "Struct:CustomLayer");
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
fn test_half_precision_with_qualified_name() {
let adapter = HalfPrecisionAdapter::new().with_module("Struct:CustomLayer");
let adapted = adapt_in(&adapter, "custom.weight", shape![5], "Struct:CustomLayer");
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
fn test_half_precision_chain() {
let adapter = PyTorchToBurnAdapter.chain(HalfPrecisionAdapter::new());
let adapted = adapt_in(&adapter, "fc.weight", shape![10, 5], module_names::LINEAR);
assert_eq!(adapted.shape, shape![5, 10]);
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
fn test_half_precision_skips_no_container() {
let adapter = HalfPrecisionAdapter::new();
let adapted = adapt_without_module(&adapter, tensor("fc.weight", shape![2, 3]));
assert_eq!(adapted.dtype, DType::F32);
}
#[test]
fn test_half_precision_skips_non_float() {
use burn_core::tensor::quantization::QuantScheme;
let adapter = HalfPrecisionAdapter::new();
let qfloat_dtype = DType::QFloat(QuantScheme::default());
let shape = shape![2, 3];
let bytes = Bytes::from_bytes_vec(vec![0u8; bridge::data_len(qfloat_dtype, &shape)]);
let qfloat = PackTensor::new("fc.weight".to_string(), qfloat_dtype, shape, None, bytes);
let containers = containers(module_names::LINEAR);
let adapted = adapter.adapt(qfloat, ModuleContext::new(&containers));
assert_eq!(adapted.dtype, qfloat_dtype);
}
#[test]
fn test_half_precision_default_module_count() {
let adapter = HalfPrecisionAdapter::new();
assert_eq!(adapter.modules.len(), 14);
}
#[test]
fn test_half_precision_without_module_qualified() {
let adapter = HalfPrecisionAdapter::new().without_module("Struct:LayerNorm");
let adapted = adapt_in(&adapter, "norm.gamma", shape![10], module_names::LAYER_NORM);
assert_eq!(adapted.dtype, DType::F32);
}
fn float_tensor(dtype: DType) -> PackTensor {
let values = vec![1.0f32, -2.0, 0.5, 4.0, -0.25, 8.0];
let data = TensorData::new(values, shape![2, 3]).convert_dtype(dtype);
tensor_with("fc.weight", data)
}
#[test]
fn test_float_cast_bf16_to_f16() {
let adapter = FloatCastAdapter::to(DType::F16);
let containers = containers(module_names::LINEAR);
let adapted = adapter.adapt(float_tensor(DType::BF16), ModuleContext::new(&containers));
assert_eq!(adapted.dtype, DType::F16);
assert_eq!(adapted.shape, shape![2, 3]);
assert_eq!(adapted.name, "fc.weight");
let data = bridge::to_data(&adapted).unwrap();
assert_eq!(data.dtype, DType::F16);
let values = data
.convert_dtype(DType::F32)
.try_into_vec::<f32>()
.unwrap();
assert_eq!(values, vec![1.0, -2.0, 0.5, 4.0, -0.25, 8.0]);
}
#[test]
fn test_float_cast_converts_any_module_type() {
let adapter = FloatCastAdapter::to(DType::F16);
let containers = containers("Struct:CustomLayer");
let adapted = adapter.adapt(float_tensor(DType::F32), ModuleContext::new(&containers));
assert_eq!(adapted.dtype, DType::F16);
let adapted = adapt_without_module(&adapter, float_tensor(DType::F32));
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
fn test_float_cast_passthrough_on_target_dtype() {
let adapter = FloatCastAdapter::to(DType::F16);
let adapted = adapt_without_module(&adapter, float_tensor(DType::F16));
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
fn test_float_cast_skips_non_float() {
let adapter = FloatCastAdapter::to(DType::F16);
let containers = containers(module_names::EMBEDDING);
let data = TensorData::new(vec![1i64, 2, 3], shape![3]);
let adapted = adapter.adapt(tensor_with("idx", data), ModuleContext::new(&containers));
assert_eq!(adapted.dtype, DType::I64);
}
#[test]
fn a_quantized_linear_weight_is_not_transposed() {
use burn_core::tensor::quantization::{QuantScheme, QuantValue};
use burn_core::tensor::{Device, Distribution, Tensor};
let device = Device::default();
let quantized = Tensor::<2>::random(shape![3, 8], Distribution::Default, &device)
.quantize_dynamic(&QuantScheme::default().with_value(QuantValue::Q4S));
let source = bridge::from_tensor(&quantized, "fc.weight".to_string(), None);
let (shape, bytes) = (
source.shape.clone(),
bridge::to_data(&source).unwrap().bytes.to_vec(),
);
let containers = containers(module_names::LINEAR);
let adapted = BurnToPyTorchAdapter.adapt(source, ModuleContext::new(&containers));
assert_eq!(
adapted.shape, shape,
"a quantized weight must pass through unchanged"
);
assert_eq!(bridge::to_data(&adapted).unwrap().bytes.to_vec(), bytes);
}
#[test]
fn module_type_is_none_without_a_user_defined_module() {
let containers = vec!["Vec".to_string(), "Array".to_string()];
assert_eq!(ModuleContext::new(&containers).module_type(), None);
}
#[test]
fn test_float_cast_chain_after_pytorch() {
let adapter = PyTorchToBurnAdapter.chain(FloatCastAdapter::to(DType::F16));
let adapted = adapt_in(&adapter, "fc.weight", shape![10, 5], module_names::LINEAR);
assert_eq!(adapted.shape, shape![5, 10]);
assert_eq!(adapted.dtype, DType::F16);
}
#[test]
#[should_panic(expected = "must be a float dtype")]
fn test_float_cast_rejects_non_float_target() {
let _ = FloatCastAdapter::to(DType::I32);
}
#[test]
fn test_half_precision_with_module_batch_norm_opt_in() {
let adapter = HalfPrecisionAdapter::new().with_module("BatchNorm");
let adapted = adapt_in(&adapter, "bn.weight", shape![10], module_names::BATCH_NORM);
assert_eq!(adapted.dtype, DType::F16);
}
}