1use 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
25pub 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
34pub 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
43pub 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 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 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}