Skip to main content

tnc/io/
hdf5.rs

1//! Import and export of tensors or tensor networks as HDF5 files.
2//!
3//! The files follow this structure:
4//! ```text
5//! tensors/
6//!     tensor: n-dimensional dataset
7//!         attrs:
8//!             - bids
9//!             - tids
10//! ```
11//!  There is a single `tensors/` group containing multiple tensor datasets. Each
12//! `tensor` is a flattened tensor with dimensions `shape`. The `tid` is the unique
13//! positive integer used to identify each tensor, with the output tensor, identified
14//! by `-1`, containing output bond dimensions and no tensor data. The `bids` are a
15//! list of integers corresponding to the bond ids of in each tensor.
16
17use std::path::Path;
18
19use hdf5_metno::{File, Result};
20use num_complex::Complex64;
21
22use crate::tensornetwork::tensor::{CompositeTensor, LeafTensor};
23use crate::tensornetwork::tensordata::{DataTensor, TensorData};
24
25/// Loads a tensor network from a HDF5 file.
26pub fn load_tensor<P>(filename: P) -> Result<CompositeTensor>
27where
28    P: AsRef<Path>,
29{
30    let file = File::open(filename)?;
31    read_tensor(&file)
32}
33
34/// Loads a single tensor from a HDF5 file.
35pub fn load_data<P>(filename: P) -> Result<DataTensor>
36where
37    P: AsRef<Path>,
38{
39    let file = File::open(filename)?;
40    read_data(&file)
41}
42
43/// Stores a single tensor in a HDF5 file.
44pub fn store_data<P>(filename: P, tensor: &DataTensor) -> Result<()>
45where
46    P: AsRef<Path>,
47{
48    let file = File::create(filename)?;
49    write_data(&file, tensor)
50}
51
52fn read_tensor(file: &File) -> Result<CompositeTensor> {
53    let gr = file.group("/tensors")?;
54    let tensor_names = gr.member_names()?;
55
56    // We don't track external legs explicitly, hence the following is commented out
57    // // Output tensor is always labelled as -1
58    // let out_tensor = gr.dataset("-1")?;
59    // let out_tensor_bids = out_tensor.attr("bids")?;
60    // let out_bond_ids = out_tensor_bids.read_1d::<usize>()?;
61
62    let mut new_tensor_network = CompositeTensor::default();
63    new_tensor_network.reserve(tensor_names.len());
64    for tensor_name in tensor_names {
65        if tensor_name == "-1" {
66            continue;
67        }
68        let tensor = gr.dataset(&tensor_name)?;
69        let bond_ids = tensor.attr("bids").unwrap().read_1d::<usize>()?;
70        let tensor_dataset = gr.dataset(&tensor_name).unwrap().read_dyn::<Complex64>()?;
71        let tensor_shape = tensor_dataset.shape();
72        let bond_dims = tensor_shape.iter().map(|s| *s as u64).collect();
73        let mut new_tensor = LeafTensor::new(bond_ids.to_vec(), bond_dims);
74        new_tensor.set_tensor_data(TensorData::Matrix(tensor_dataset));
75        new_tensor_network.push_tensor(new_tensor);
76    }
77
78    Ok(new_tensor_network)
79}
80
81fn read_data(file: &File) -> Result<DataTensor> {
82    let gr = file.group("/tensors")?;
83    let tensor_name = gr.member_names()?;
84
85    let tensor_dataset = gr
86        .dataset(&tensor_name[0])
87        .unwrap()
88        .read_dyn::<Complex64>()?;
89    Ok(tensor_dataset)
90}
91
92fn write_data(file: &File, tensor: &DataTensor) -> Result<()> {
93    let gr = file.create_group("/tensors")?;
94    let tensor_dataset = gr.new_dataset_builder().with_data(tensor);
95    tensor_dataset.create("-1")?;
96    file.flush()
97}
98
99#[cfg(test)]
100mod tests {
101    use super::*;
102
103    use std::iter::zip;
104
105    use approx::assert_abs_diff_eq;
106    use hdf5_metno::{AttributeBuilder, File, Result};
107    use ndarray::{array, Array2};
108    use num_complex::Complex64;
109    use rand::{
110        distr::{Alphanumeric, SampleString},
111        rng,
112    };
113
114    use crate::tensornetwork::tensordata::TensorData;
115
116    /// Creates a new HDF5 file in memory.
117    /// This method is taken from the hdf5 crate integration tests:
118    /// <https://github.com/aldanor/hdf5-rust/blob/694e900972fbf5ffbdd1a2294f57a2cc3a91c994/hdf5/tests/common/util.rs#L7>.
119    fn new_in_memory_file() -> Result<File> {
120        let random_filename = Alphanumeric.sample_string(&mut rng(), 8);
121        File::with_options()
122            .with_access_plist(|p| p.core_filebacked(false))
123            .create(random_filename)
124    }
125
126    fn create_hdf5_tensor() -> Result<File> {
127        let new_file = new_in_memory_file()?;
128        let tensor_group = new_file.create_group("./tensors")?;
129        let dataset_builder = tensor_group.new_dataset_builder();
130        let dataset = dataset_builder.empty::<Complex64>().create("-1")?;
131        let attribute = AttributeBuilder::new(&dataset);
132        let bid = array![0, 1];
133        let attribute = attribute.with_data(&bid);
134        attribute.create("bids")?;
135
136        let data = Array2::<Complex64>::from_shape_vec(
137            (2, 2),
138            vec![
139                Complex64::new(1.0, 0.0),
140                Complex64::new(0.0, 2.0),
141                Complex64::new(3.0, 0.0),
142                Complex64::new(0.0, 1.0),
143            ],
144        )?
145        .into_dyn();
146        let dataset_builder2 = tensor_group.new_dataset_builder();
147        let dataset_data_builder2 = dataset_builder2.with_data(&data);
148        let dataset2 = dataset_data_builder2.create("0")?;
149        let attribute2 = AttributeBuilder::new(&dataset2);
150        let bid2 = array![0, 1];
151        let attribute2 = attribute2.with_data(&bid2);
152        attribute2.create("bids")?;
153
154        new_file.flush()?;
155        Ok(new_file)
156    }
157
158    fn create_hdf5_data() -> Result<File> {
159        let new_file = new_in_memory_file()?;
160        let tensor_group = new_file.create_group("./tensors")?;
161        let dataset_builder = tensor_group.new_dataset_builder();
162        let data = Array2::<Complex64>::from_shape_vec(
163            (2, 2),
164            vec![
165                Complex64::new(1.0, 0.0),
166                Complex64::new(0.0, 2.0),
167                Complex64::new(3.0, 0.0),
168                Complex64::new(0.0, 1.0),
169            ],
170        )?
171        .into_dyn();
172        let dataset_data_builder = dataset_builder.with_data(&data);
173        dataset_data_builder.create("-1")?;
174        new_file.flush()?;
175        Ok(new_file)
176    }
177
178    #[test]
179    fn test_load_data() {
180        let file = create_hdf5_data().unwrap();
181        let tensor_data = read_data(&file).unwrap();
182
183        let ref_data = array![
184            Complex64::new(1.0, 0.0),
185            Complex64::new(0.0, 2.0),
186            Complex64::new(3.0, 0.0),
187            Complex64::new(0.0, 1.0),
188        ];
189        for (u, v) in zip(ref_data.iter(), tensor_data.flatten().iter()) {
190            assert_abs_diff_eq!(u.re, v.re, epsilon = 1e-8);
191            assert_abs_diff_eq!(u.im, v.im, epsilon = 1e-8);
192        }
193    }
194
195    #[test]
196    fn test_load_tensor() {
197        let file = create_hdf5_tensor().unwrap();
198        let tensor = read_tensor(&file).unwrap();
199
200        let mut ref_tensor = LeafTensor::new(vec![0, 1], vec![2, 2]);
201        ref_tensor.set_tensor_data(TensorData::new_from_data(
202            &[2, 2],
203            vec![
204                Complex64::new(1.0, 0.0),
205                Complex64::new(0.0, 2.0),
206                Complex64::new(3.0, 0.0),
207                Complex64::new(0.0, 1.0),
208            ],
209        ));
210        let ref_tn = CompositeTensor::new(vec![ref_tensor]);
211        assert_abs_diff_eq!(&tensor, &ref_tn);
212    }
213
214    #[test]
215    fn test_write_read() {
216        let file = new_in_memory_file().unwrap();
217        let data = vec![
218            Complex64::new(1.0, 0.0),
219            Complex64::new(0.0, -2.0),
220            Complex64::new(-3.0, 0.0),
221            Complex64::new(-2.0, -1.0),
222            Complex64::new(0.0, 0.0),
223            Complex64::new(0.5, 2.0),
224        ];
225        let tensor = Array2::<Complex64>::from_shape_vec((2, 3), data)
226            .unwrap()
227            .into_dyn();
228
229        write_data(&file, &tensor).unwrap();
230        let read = read_data(&file).unwrap();
231
232        assert_abs_diff_eq!(&tensor, &read);
233    }
234}