use crate::TensorSnapshot;
use alloc::boxed::Box;
use alloc::rc::Rc;
use alloc::string::String;
use alloc::string::ToString;
use alloc::vec;
use burn_tensor::TensorData;
mod module_names {
#[allow(unused_imports)]
use burn_nn::{BatchNorm, GroupNorm, LayerNorm, Linear};
pub const LINEAR: &str = "Linear";
pub const BATCH_NORM: &str = "BatchNorm";
pub const LAYER_NORM: &str = "LayerNorm";
pub const GROUP_NORM: &str = "GroupNorm";
}
pub trait ModuleAdapter: Send + Sync {
fn adapt(&self, snapshot: &TensorSnapshot) -> TensorSnapshot;
fn clone_box(&self) -> Box<dyn ModuleAdapter>;
}
impl Clone for Box<dyn ModuleAdapter> {
fn clone(&self) -> Self {
self.clone_box()
}
}
#[derive(Debug, Clone, Default)]
pub struct IdentityAdapter;
impl ModuleAdapter for IdentityAdapter {
fn adapt(&self, snapshot: &TensorSnapshot) -> TensorSnapshot {
snapshot.clone()
}
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, snapshot: &TensorSnapshot) -> TensorSnapshot {
adapt_pytorch_tensor(snapshot, PyTorchConversionDirection::PyTorchToBurn)
}
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, snapshot: &TensorSnapshot) -> TensorSnapshot {
adapt_pytorch_tensor(snapshot, PyTorchConversionDirection::BurnToPyTorch)
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
#[derive(Debug, Clone, Copy)]
enum PyTorchConversionDirection {
PyTorchToBurn,
BurnToPyTorch,
}
fn adapt_pytorch_tensor(
snapshot: &TensorSnapshot,
direction: PyTorchConversionDirection,
) -> TensorSnapshot {
let (path_stack, param_name) = match get_path_and_param(snapshot) {
Some(result) => result,
None => return snapshot.clone(),
};
let container_type = match snapshot.container_stack.as_ref().and_then(|s| s.last()) {
Some(ct) => ct,
None => return snapshot.clone(),
};
match container_type.as_str() {
module_names::LINEAR if param_name == "weight" && snapshot.shape.len() == 2 => {
transpose_2d_tensor(snapshot)
}
module_names::BATCH_NORM | module_names::LAYER_NORM | module_names::GROUP_NORM => {
let new_name = match direction {
PyTorchConversionDirection::PyTorchToBurn => match param_name {
"weight" => "gamma",
"bias" => "beta",
_ => return snapshot.clone(),
},
PyTorchConversionDirection::BurnToPyTorch => match param_name {
"gamma" => "weight",
"beta" => "bias",
_ => return snapshot.clone(),
},
};
rename_parameter(snapshot, path_stack, new_name)
}
_ => snapshot.clone(),
}
}
fn get_path_and_param(snapshot: &TensorSnapshot) -> Option<(&[String], &str)> {
let path_stack = snapshot.path_stack.as_ref()?;
let param_name = path_stack.last()?.as_str();
Some((path_stack.as_slice(), param_name))
}
fn rename_parameter(
snapshot: &TensorSnapshot,
path_stack: &[String],
new_name: &str,
) -> TensorSnapshot {
let mut new_path = path_stack.to_vec();
*new_path.last_mut().unwrap() = new_name.to_string();
TensorSnapshot::from_closure(
snapshot.clone_data_fn(),
snapshot.dtype,
snapshot.shape.clone(),
new_path,
snapshot.container_stack.clone().unwrap_or_default(),
snapshot.tensor_id.unwrap_or_default(),
)
}
fn transpose_2d_tensor(snapshot: &TensorSnapshot) -> TensorSnapshot {
if snapshot.shape.len() != 2 {
return snapshot.clone();
}
let original_data_fn = snapshot.clone_data_fn();
let dtype = snapshot.dtype;
let transposed_shape = vec![snapshot.shape[1], snapshot.shape[0]];
let transposed_data_fn = Rc::new(move || {
let data = original_data_fn()?;
Ok(transpose_tensor_data(data))
});
TensorSnapshot::from_closure(
transposed_data_fn,
dtype,
transposed_shape,
snapshot.path_stack.clone().unwrap_or_default(),
snapshot.container_stack.clone().unwrap_or_default(),
snapshot.tensor_id.unwrap_or_default(),
)
}
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::rc::Rc;
use burn_tensor::{DType, TensorData};
fn create_test_snapshot(path: &str, shape: Vec<usize>, container_type: &str) -> TensorSnapshot {
let path_parts: Vec<String> = path.split('.').map(|s| s.to_string()).collect();
let values = vec![1.0f32; shape.iter().product()];
let data = TensorData::new(values, shape.clone());
TensorSnapshot::from_closure(
Rc::new(move || Ok(data.clone())),
DType::F32,
shape,
path_parts,
vec![container_type.to_string()],
burn_core::module::ParamId::new(),
)
}
#[test]
fn test_pytorch_to_burn_linear_weight() {
let adapter = PyTorchToBurnAdapter;
let snapshot = create_test_snapshot("fc.weight", vec![10, 5], module_names::LINEAR);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.shape, vec![5, 10]);
let snapshot = create_test_snapshot("fc.bias", vec![10], module_names::LINEAR);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.shape, vec![10]);
}
#[test]
fn test_pytorch_to_burn_norm_params() {
let adapter = PyTorchToBurnAdapter;
let snapshot = create_test_snapshot("norm.weight", vec![10], module_names::BATCH_NORM);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.full_path(), "norm.gamma");
let snapshot = create_test_snapshot("norm.bias", vec![10], module_names::BATCH_NORM);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.full_path(), "norm.beta");
}
#[test]
fn test_burn_to_pytorch_linear_weight() {
let adapter = BurnToPyTorchAdapter;
let snapshot = create_test_snapshot("fc.weight", vec![5, 10], module_names::LINEAR);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.shape, vec![10, 5]);
}
#[test]
fn test_burn_to_pytorch_norm_params() {
let adapter = BurnToPyTorchAdapter;
let snapshot = create_test_snapshot("norm.gamma", vec![10], module_names::BATCH_NORM);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.full_path(), "norm.weight");
let snapshot = create_test_snapshot("norm.beta", vec![10], module_names::BATCH_NORM);
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.full_path(), "norm.bias");
}
#[test]
fn test_transpose_different_dtypes() {
let f32_data = TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]);
let transposed = transpose_tensor_data(f32_data);
assert_eq!(transposed.shape, vec![3, 2]);
let values = transposed.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], vec![2, 3]);
let transposed = transpose_tensor_data(i32_data);
assert_eq!(transposed.shape, vec![3, 2]);
let values = transposed.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], vec![2, 2]);
let transposed = transpose_tensor_data(f64_data);
assert_eq!(transposed.shape, vec![2, 2]);
let values = transposed.to_vec::<f64>().unwrap();
assert_eq!(values, vec![1.0, 3.0, 2.0, 4.0]);
}
#[test]
fn test_no_container_info() {
let adapter = PyTorchToBurnAdapter;
let mut snapshot = create_test_snapshot("fc.weight", vec![10, 5], module_names::LINEAR);
snapshot.container_stack = None;
let adapted = adapter.adapt(&snapshot);
assert_eq!(adapted.shape, vec![10, 5]);
let mut snapshot2 = create_test_snapshot("other.weight", vec![10, 5], "Other");
snapshot2.container_stack = None;
let adapted2 = adapter.adapt(&snapshot2);
assert_eq!(adapted2.shape, vec![10, 5]); }
}