1use std::collections::{BTreeSet, HashMap};
2use std::fs::{self, OpenOptions};
3use std::path::Path;
4
5use half::f16;
6use memmap2::MmapOptions;
7
8use crate::any_trx_file::{AnyTrxFile, PositionsRef};
9use crate::dtype::{DType, TrxScalar};
10use crate::error::{Result, TrxError};
11use crate::mmap_backing::MmapBacking;
12use crate::trx_file::{DataArray, DataArrayInfo, FromF32, TrxFile, TrxParts};
13
14#[derive(Clone, Debug, Default)]
16pub struct ConcatenateOptions {
17 pub delete_dpv: bool,
19 pub delete_dps: bool,
21 pub delete_groups: bool,
23 pub positions_dtype: Option<DType>,
25 pub input_group_names: Vec<Option<String>>,
30}
31
32#[derive(Clone, Copy, Debug)]
33struct ArraySpec {
34 ncols: usize,
35 dtype: DType,
36}
37
38#[derive(Clone, Copy, Debug)]
39struct GroupSpec {
40 dtype: DType,
41 len: usize,
42}
43
44#[derive(Clone, Debug)]
45struct InputGroupRename {
46 aggregate_name: Option<String>,
47 renamed_groups: Vec<(String, String)>,
48}
49
50pub fn concatenate_any_trx(
55 inputs: &[&AnyTrxFile],
56 options: &ConcatenateOptions,
57) -> Result<AnyTrxFile> {
58 if inputs.is_empty() {
59 return Err(TrxError::Argument("no shards to merge".into()));
60 }
61
62 let first = inputs[0];
63 let target_dtype = options.positions_dtype.unwrap_or_else(|| first.dtype());
64 let total_streamlines = inputs.iter().map(|input| input.nb_streamlines()).sum();
65 let total_vertices = inputs.iter().map(|input| input.nb_vertices()).sum();
66 let group_renames = normalized_group_names(inputs.len(), &options.input_group_names)?
67 .into_iter()
68 .zip(inputs.iter())
69 .map(|(group_name, input)| {
70 effective_group_rename(group_name, group_map(input).keys().cloned())
71 })
72 .collect::<Vec<_>>();
73
74 let dps_specs = retained_array_specs(inputs, ArrayKind::Dps, options.delete_dps)?;
75 let dpv_specs = retained_array_specs(inputs, ArrayKind::Dpv, options.delete_dpv)?;
76 let group_specs = retained_group_specs_with_renames(inputs, &group_renames)?;
77
78 match target_dtype {
79 DType::Float16 => concatenate_into::<f16>(
80 inputs,
81 total_streamlines,
82 total_vertices,
83 dps_specs,
84 dpv_specs,
85 group_specs,
86 &group_renames,
87 options,
88 )
89 .map(AnyTrxFile::F16),
90 DType::Float32 => concatenate_into::<f32>(
91 inputs,
92 total_streamlines,
93 total_vertices,
94 dps_specs,
95 dpv_specs,
96 group_specs,
97 &group_renames,
98 options,
99 )
100 .map(AnyTrxFile::F32),
101 DType::Float64 => concatenate_into::<f64>(
102 inputs,
103 total_streamlines,
104 total_vertices,
105 dps_specs,
106 dpv_specs,
107 group_specs,
108 &group_renames,
109 options,
110 )
111 .map(AnyTrxFile::F64),
112 other => Err(TrxError::DType(format!(
113 "TRX positions must be float16, float32, or float64, got {other}"
114 ))),
115 }
116}
117
118pub fn merge_trx_shards<P: TrxScalar>(shards: &[&TrxFile<P>]) -> Result<TrxFile<P>> {
120 if shards.is_empty() {
121 return Err(TrxError::Argument("no shards to merge".into()));
122 }
123
124 let total_streamlines = shards.iter().map(|input| input.nb_streamlines()).sum();
125 let total_vertices = shards.iter().map(|input| input.nb_vertices()).sum();
126 let dps_specs = retained_array_specs_typed(shards, ArrayKind::Dps, false)?;
127 let dpv_specs = retained_array_specs_typed(shards, ArrayKind::Dpv, false)?;
128 let group_specs = retained_group_specs_typed(shards)?;
129
130 concatenate_typed_from_trx(
131 shards,
132 total_streamlines,
133 total_vertices,
134 dps_specs,
135 dpv_specs,
136 group_specs,
137 &ConcatenateOptions {
138 positions_dtype: Some(P::DTYPE),
139 ..Default::default()
140 },
141 )
142}
143
144#[allow(clippy::too_many_arguments)]
145fn concatenate_into<P>(
146 inputs: &[&AnyTrxFile],
147 total_streamlines: usize,
148 total_vertices: usize,
149 dps_specs: HashMap<String, ArraySpec>,
150 dpv_specs: HashMap<String, ArraySpec>,
151 group_specs: HashMap<String, GroupSpec>,
152 group_renames: &[InputGroupRename],
153 options: &ConcatenateOptions,
154) -> Result<TrxFile<P>>
155where
156 P: TrxScalar + FromF32,
157{
158 let mut header = inputs[0].header().clone();
159 header.nb_streamlines = total_streamlines as u64;
160 header.nb_vertices = total_vertices as u64;
161
162 let tempdir = tempfile::TempDir::new()?;
163 let tempdir_path = tempdir.path().to_path_buf();
164
165 let mut positions_backing = create_mmap_backing(
166 &tempdir_path.join(format!("positions.3.{}", P::DTYPE.name())),
167 total_vertices * 3 * std::mem::size_of::<P>(),
168 )?;
169 let mut offsets_backing = create_mmap_backing(
170 &tempdir_path.join("offsets.uint32"),
171 (total_streamlines + 1) * 4,
172 )?;
173 let mut dps = create_data_map(&tempdir_path.join("dps"), &dps_specs, total_streamlines)?;
174 let mut dpv = create_data_map(&tempdir_path.join("dpv"), &dpv_specs, total_vertices)?;
175 let mut groups = if options.delete_groups {
176 HashMap::new()
177 } else {
178 create_group_map(&tempdir_path.join("groups"), &group_specs)?
179 };
180
181 fill_positions::<P>(inputs, &mut positions_backing)?;
182 fill_offsets(inputs, &mut offsets_backing)?;
183 copy_retained_arrays(inputs, &dps_specs, &mut dps, ArrayKind::Dps)?;
184 copy_retained_arrays(inputs, &dpv_specs, &mut dpv, ArrayKind::Dpv)?;
185 if !options.delete_groups {
186 copy_groups(inputs, group_renames, &group_specs, &mut groups)?;
187 }
188
189 Ok(TrxFile::from_parts(TrxParts {
190 header,
191 positions_backing,
192 offsets_backing,
193 dps,
194 dpv,
195 groups,
196 dpg: HashMap::new(),
197 tempdir: Some(tempdir),
198 }))
199}
200
201fn concatenate_typed_from_trx<P>(
202 inputs: &[&TrxFile<P>],
203 total_streamlines: usize,
204 total_vertices: usize,
205 dps_specs: HashMap<String, ArraySpec>,
206 dpv_specs: HashMap<String, ArraySpec>,
207 group_specs: HashMap<String, GroupSpec>,
208 options: &ConcatenateOptions,
209) -> Result<TrxFile<P>>
210where
211 P: TrxScalar,
212{
213 let mut header = inputs[0].header().clone();
214 header.nb_streamlines = total_streamlines as u64;
215 header.nb_vertices = total_vertices as u64;
216
217 let tempdir = tempfile::TempDir::new()?;
218 let tempdir_path = tempdir.path().to_path_buf();
219
220 let mut positions_backing = create_mmap_backing(
221 &tempdir_path.join(format!("positions.3.{}", P::DTYPE.name())),
222 total_vertices * 3 * std::mem::size_of::<P>(),
223 )?;
224 let mut offsets_backing = create_mmap_backing(
225 &tempdir_path.join("offsets.uint32"),
226 (total_streamlines + 1) * 4,
227 )?;
228 let mut dps = create_data_map(&tempdir_path.join("dps"), &dps_specs, total_streamlines)?;
229 let mut dpv = create_data_map(&tempdir_path.join("dpv"), &dpv_specs, total_vertices)?;
230 let mut groups = if options.delete_groups {
231 HashMap::new()
232 } else {
233 create_group_map(&tempdir_path.join("groups"), &group_specs)?
234 };
235
236 fill_positions_typed(inputs, &mut positions_backing)?;
237 fill_offsets_typed(inputs, &mut offsets_backing)?;
238 copy_retained_arrays_typed(inputs, &dps_specs, &mut dps, ArrayKind::Dps)?;
239 copy_retained_arrays_typed(inputs, &dpv_specs, &mut dpv, ArrayKind::Dpv)?;
240 if !options.delete_groups {
241 copy_groups_typed(inputs, &group_specs, &mut groups)?;
242 }
243
244 Ok(TrxFile::from_parts(TrxParts {
245 header,
246 positions_backing,
247 offsets_backing,
248 dps,
249 dpv,
250 groups,
251 dpg: HashMap::new(),
252 tempdir: Some(tempdir),
253 }))
254}
255
256#[derive(Clone, Copy)]
257enum ArrayKind {
258 Dps,
259 Dpv,
260}
261
262fn retained_array_specs(
263 inputs: &[&AnyTrxFile],
264 kind: ArrayKind,
265 delete_all: bool,
266) -> Result<HashMap<String, ArraySpec>> {
267 let union_keys = array_keys_union(inputs, kind);
268 let reference = array_infos(inputs[0], kind);
269
270 for input in inputs.iter().skip(1) {
271 let current = array_infos(input, kind);
272 for key in &union_keys {
273 match (reference.get(key), current.get(key)) {
274 (Some(lhs), Some(rhs)) => ensure_matching_array_info(kind, key, *lhs, *rhs)?,
275 _ if delete_all => {}
276 _ => {
277 return Err(TrxError::Argument(format!(
278 "{} key '{key}' does not exist in all TrxFile inputs",
279 array_kind_name(kind)
280 )))
281 }
282 }
283 }
284 }
285
286 if delete_all {
287 return Ok(HashMap::new());
288 }
289
290 Ok(reference
291 .into_iter()
292 .map(|(name, info)| {
293 (
294 name,
295 ArraySpec {
296 ncols: info.ncols,
297 dtype: info.dtype,
298 },
299 )
300 })
301 .collect())
302}
303
304fn retained_group_specs_with_renames(
305 inputs: &[&AnyTrxFile],
306 group_renames: &[InputGroupRename],
307) -> Result<HashMap<String, GroupSpec>> {
308 let mut specs: HashMap<String, GroupSpec> = HashMap::new();
309 let mut lengths: HashMap<String, usize> = HashMap::new();
310
311 for (input, rename) in inputs.iter().zip(group_renames.iter()) {
312 let groups = group_infos(input);
313 if groups.is_empty() {
314 if let Some(name) = &rename.aggregate_name {
315 match specs.get(name) {
316 Some(existing) if existing.dtype != DType::UInt32 => {
317 return Err(TrxError::Argument(format!(
318 "group key '{name}' has different dtypes across inputs"
319 )))
320 }
321 Some(_) => {}
322 None => {
323 specs.insert(
324 name.clone(),
325 GroupSpec {
326 dtype: DType::UInt32,
327 len: 0,
328 },
329 );
330 }
331 }
332 *lengths.entry(name.clone()).or_insert(0usize) += input.nb_streamlines();
333 }
334 continue;
335 }
336
337 for (source_name, info) in groups {
338 let name = rename
339 .renamed_groups
340 .iter()
341 .find_map(|(from, to)| (from == &source_name).then_some(to))
342 .cloned()
343 .unwrap_or(source_name);
344 match specs.get(&name) {
345 Some(existing) if existing.dtype != info.dtype => {
346 return Err(TrxError::Argument(format!(
347 "group key '{name}' has different dtypes across inputs"
348 )))
349 }
350 Some(_) => {}
351 None => {
352 specs.insert(
353 name.clone(),
354 GroupSpec {
355 dtype: info.dtype,
356 len: 0,
357 },
358 );
359 }
360 }
361 *lengths.entry(name).or_insert(0usize) += info.nrows;
362 }
363 }
364
365 for (name, len) in lengths {
366 specs
367 .get_mut(&name)
368 .expect("lengths and specs must be in sync")
369 .len = len;
370 }
371
372 Ok(specs)
373}
374
375fn retained_array_specs_typed<P: TrxScalar>(
376 inputs: &[&TrxFile<P>],
377 kind: ArrayKind,
378 delete_all: bool,
379) -> Result<HashMap<String, ArraySpec>> {
380 let union_keys: BTreeSet<String> = inputs
381 .iter()
382 .flat_map(|input| match kind {
383 ArrayKind::Dps => input.dps_arrays().keys(),
384 ArrayKind::Dpv => input.dpv_arrays().keys(),
385 })
386 .cloned()
387 .collect();
388 let reference: HashMap<String, DataArrayInfo> = match kind {
389 ArrayKind::Dps => inputs[0]
390 .dps_arrays()
391 .iter()
392 .map(|(name, arr)| (name.clone(), arr.info()))
393 .collect(),
394 ArrayKind::Dpv => inputs[0]
395 .dpv_arrays()
396 .iter()
397 .map(|(name, arr)| (name.clone(), arr.info()))
398 .collect(),
399 };
400
401 for input in inputs.iter().skip(1) {
402 let current: HashMap<String, DataArrayInfo> = match kind {
403 ArrayKind::Dps => input
404 .dps_arrays()
405 .iter()
406 .map(|(name, arr)| (name.clone(), arr.info()))
407 .collect(),
408 ArrayKind::Dpv => input
409 .dpv_arrays()
410 .iter()
411 .map(|(name, arr)| (name.clone(), arr.info()))
412 .collect(),
413 };
414 for key in &union_keys {
415 match (reference.get(key), current.get(key)) {
416 (Some(lhs), Some(rhs)) => ensure_matching_array_info(kind, key, *lhs, *rhs)?,
417 _ if delete_all => {}
418 _ => {
419 return Err(TrxError::Argument(format!(
420 "{} key '{key}' does not exist in all TrxFile inputs",
421 array_kind_name(kind)
422 )))
423 }
424 }
425 }
426 }
427
428 if delete_all {
429 return Ok(HashMap::new());
430 }
431
432 Ok(reference
433 .into_iter()
434 .map(|(name, info)| {
435 (
436 name,
437 ArraySpec {
438 ncols: info.ncols,
439 dtype: info.dtype,
440 },
441 )
442 })
443 .collect())
444}
445
446fn retained_group_specs_typed<P: TrxScalar>(
447 inputs: &[&TrxFile<P>],
448) -> Result<HashMap<String, GroupSpec>> {
449 let mut specs: HashMap<String, GroupSpec> = HashMap::new();
450 let mut lengths: HashMap<String, usize> = HashMap::new();
451
452 for input in inputs {
453 for (name, arr) in input.group_arrays() {
454 match specs.get(name) {
455 Some(existing) if existing.dtype != arr.dtype() => {
456 return Err(TrxError::Argument(format!(
457 "group key '{name}' has different dtypes across inputs"
458 )))
459 }
460 Some(_) => {}
461 None => {
462 specs.insert(
463 name.clone(),
464 GroupSpec {
465 dtype: arr.dtype(),
466 len: 0,
467 },
468 );
469 }
470 }
471 *lengths.entry(name.clone()).or_insert(0usize) += arr.nrows();
472 }
473 }
474
475 for (name, len) in lengths {
476 specs
477 .get_mut(&name)
478 .expect("lengths and specs must be in sync")
479 .len = len;
480 }
481 Ok(specs)
482}
483
484fn array_keys_union(inputs: &[&AnyTrxFile], kind: ArrayKind) -> BTreeSet<String> {
485 let mut keys = BTreeSet::new();
486 for input in inputs {
487 for key in array_infos(input, kind).into_keys() {
488 keys.insert(key);
489 }
490 }
491 keys
492}
493
494fn array_infos(file: &AnyTrxFile, kind: ArrayKind) -> HashMap<String, DataArrayInfo> {
495 match kind {
496 ArrayKind::Dps => file.dps_entries().into_iter().collect(),
497 ArrayKind::Dpv => file.dpv_entries().into_iter().collect(),
498 }
499}
500
501fn group_infos(file: &AnyTrxFile) -> HashMap<String, DataArrayInfo> {
502 file.with_typed(
503 |trx| {
504 trx.group_arrays()
505 .iter()
506 .map(|(name, arr)| (name.clone(), arr.info()))
507 .collect()
508 },
509 |trx| {
510 trx.group_arrays()
511 .iter()
512 .map(|(name, arr)| (name.clone(), arr.info()))
513 .collect()
514 },
515 |trx| {
516 trx.group_arrays()
517 .iter()
518 .map(|(name, arr)| (name.clone(), arr.info()))
519 .collect()
520 },
521 )
522}
523
524fn ensure_matching_array_info(
525 kind: ArrayKind,
526 key: &str,
527 lhs: DataArrayInfo,
528 rhs: DataArrayInfo,
529) -> Result<()> {
530 if lhs.dtype != rhs.dtype {
531 return Err(TrxError::Argument(format!(
532 "{} key '{key}' has different dtypes across inputs",
533 array_kind_name(kind)
534 )));
535 }
536 if lhs.ncols != rhs.ncols {
537 return Err(TrxError::Argument(format!(
538 "{} key '{key}' has different column counts across inputs",
539 array_kind_name(kind)
540 )));
541 }
542 Ok(())
543}
544
545fn array_kind_name(kind: ArrayKind) -> &'static str {
546 match kind {
547 ArrayKind::Dps => "dps",
548 ArrayKind::Dpv => "dpv",
549 }
550}
551
552fn create_mmap_backing(path: &Path, len: usize) -> Result<MmapBacking> {
553 if len == 0 {
554 return Ok(MmapBacking::Owned(Vec::new()));
555 }
556 if let Some(parent) = path.parent() {
557 fs::create_dir_all(parent)?;
558 }
559 let file = OpenOptions::new()
560 .read(true)
561 .write(true)
562 .create(true)
563 .truncate(true)
564 .open(path)?;
565 file.set_len(len as u64)?;
566 let mmap = unsafe { MmapOptions::new().len(len).map_mut(&file)? };
567 Ok(MmapBacking::ReadWrite(mmap))
568}
569
570fn create_data_map(
571 dir: &Path,
572 specs: &HashMap<String, ArraySpec>,
573 rows: usize,
574) -> Result<HashMap<String, DataArray>> {
575 let mut out = HashMap::new();
576 if specs.is_empty() {
577 return Ok(out);
578 }
579 fs::create_dir_all(dir)?;
580 for (name, spec) in specs {
581 let filename = crate::io::filename::TrxFilename {
582 name: name.clone(),
583 ncols: spec.ncols,
584 dtype: spec.dtype,
585 }
586 .to_filename();
587 let len = rows
588 .checked_mul(spec.ncols)
589 .and_then(|v| v.checked_mul(spec.dtype.size_of()))
590 .ok_or_else(|| TrxError::Argument(format!("array '{name}' is too large")))?;
591 out.insert(
592 name.clone(),
593 DataArray::from_backing(
594 create_mmap_backing(&dir.join(filename), len)?,
595 spec.ncols,
596 spec.dtype,
597 ),
598 );
599 }
600 Ok(out)
601}
602
603fn create_group_map(
604 dir: &Path,
605 specs: &HashMap<String, GroupSpec>,
606) -> Result<HashMap<String, DataArray>> {
607 let mut out = HashMap::new();
608 if specs.is_empty() {
609 return Ok(out);
610 }
611 fs::create_dir_all(dir)?;
612 for (name, spec) in specs {
613 let filename = crate::io::filename::TrxFilename {
614 name: name.clone(),
615 ncols: 1,
616 dtype: spec.dtype,
617 }
618 .to_filename();
619 let len = spec
620 .len
621 .checked_mul(spec.dtype.size_of())
622 .ok_or_else(|| TrxError::Argument(format!("group '{name}' is too large")))?;
623 out.insert(
624 name.clone(),
625 DataArray::from_backing(
626 create_mmap_backing(&dir.join(filename), len)?,
627 1,
628 spec.dtype,
629 ),
630 );
631 }
632 Ok(out)
633}
634
635fn fill_positions<P>(inputs: &[&AnyTrxFile], backing: &mut MmapBacking) -> Result<()>
636where
637 P: TrxScalar + FromF32,
638{
639 let dst: &mut [[P; 3]] = backing.cast_slice_mut()?;
640 let mut cursor = 0usize;
641 for input in inputs {
642 let count = input.nb_vertices();
643 let target = &mut dst[cursor..cursor + count];
644 match input.positions_ref() {
645 PositionsRef::F16(src) => copy_positions(src, target),
646 PositionsRef::F32(src) => copy_positions(src, target),
647 PositionsRef::F64(src) => copy_positions(src, target),
648 }
649 cursor += count;
650 }
651 Ok(())
652}
653
654fn copy_positions<Src, Dst>(src: &[[Src; 3]], dst: &mut [[Dst; 3]])
655where
656 Src: TrxScalar,
657 Dst: TrxScalar + FromF32,
658{
659 for (src_row, dst_row) in src.iter().zip(dst.iter_mut()) {
660 dst_row[0] = Dst::from_f32(src_row[0].to_f32());
661 dst_row[1] = Dst::from_f32(src_row[1].to_f32());
662 dst_row[2] = Dst::from_f32(src_row[2].to_f32());
663 }
664}
665
666fn fill_offsets(inputs: &[&AnyTrxFile], backing: &mut MmapBacking) -> Result<()> {
667 let dst: &mut [u32] = backing.cast_slice_mut()?;
668 let mut cursor = 0usize;
669 let mut vertex_base = 0u32;
670 for input in inputs {
671 for offset in input.offsets_vec().into_iter().take(input.nb_streamlines()) {
672 dst[cursor] = offset
673 .checked_add(vertex_base)
674 .ok_or_else(|| TrxError::Argument("offset overflow during concatenate".into()))?;
675 cursor += 1;
676 }
677 vertex_base = vertex_base
678 .checked_add(input.nb_vertices() as u32)
679 .ok_or_else(|| TrxError::Argument("vertex count overflow during concatenate".into()))?;
680 }
681 dst[cursor] = vertex_base;
682 Ok(())
683}
684
685fn fill_positions_typed<P>(inputs: &[&TrxFile<P>], backing: &mut MmapBacking) -> Result<()>
686where
687 P: TrxScalar,
688{
689 let dst: &mut [[P; 3]] = backing.cast_slice_mut()?;
690 let mut cursor = 0usize;
691 for input in inputs {
692 let count = input.nb_vertices();
693 dst[cursor..cursor + count].copy_from_slice(input.positions());
694 cursor += count;
695 }
696 Ok(())
697}
698
699fn fill_offsets_typed<P: TrxScalar>(
700 inputs: &[&TrxFile<P>],
701 backing: &mut MmapBacking,
702) -> Result<()> {
703 let dst: &mut [u32] = backing.cast_slice_mut()?;
704 let mut cursor = 0usize;
705 let mut vertex_base = 0u32;
706 for input in inputs {
707 for &offset in input.offsets().iter().take(input.nb_streamlines()) {
708 dst[cursor] = offset
709 .checked_add(vertex_base)
710 .ok_or_else(|| TrxError::Argument("offset overflow during concatenate".into()))?;
711 cursor += 1;
712 }
713 vertex_base = vertex_base
714 .checked_add(input.nb_vertices() as u32)
715 .ok_or_else(|| TrxError::Argument("vertex count overflow during concatenate".into()))?;
716 }
717 dst[cursor] = vertex_base;
718 Ok(())
719}
720
721fn copy_retained_arrays(
722 inputs: &[&AnyTrxFile],
723 specs: &HashMap<String, ArraySpec>,
724 outputs: &mut HashMap<String, DataArray>,
725 kind: ArrayKind,
726) -> Result<()> {
727 for (name, spec) in specs {
728 let row_bytes = spec.ncols * spec.dtype.size_of();
729 let dst = outputs
730 .get_mut(name)
731 .expect("output map should contain every retained array")
732 .as_bytes_mut()?;
733 let mut byte_cursor = 0usize;
734 for input in inputs {
735 let src = array_map(input, kind)
736 .get(name)
737 .expect("retained array must exist in every input");
738 let bytes = src.as_bytes();
739 let len = src
740 .nrows()
741 .checked_mul(row_bytes)
742 .ok_or_else(|| TrxError::Argument(format!("array '{name}' is too large")))?;
743 dst[byte_cursor..byte_cursor + len].copy_from_slice(bytes);
744 byte_cursor += len;
745 }
746 }
747 Ok(())
748}
749
750fn copy_retained_arrays_typed<P: TrxScalar>(
751 inputs: &[&TrxFile<P>],
752 specs: &HashMap<String, ArraySpec>,
753 outputs: &mut HashMap<String, DataArray>,
754 kind: ArrayKind,
755) -> Result<()> {
756 for (name, spec) in specs {
757 let row_bytes = spec.ncols * spec.dtype.size_of();
758 let dst = outputs
759 .get_mut(name)
760 .expect("output map should contain every retained array")
761 .as_bytes_mut()?;
762 let mut byte_cursor = 0usize;
763 for input in inputs {
764 let src = match kind {
765 ArrayKind::Dps => input.dps_arrays().get(name),
766 ArrayKind::Dpv => input.dpv_arrays().get(name),
767 }
768 .expect("retained array must exist in every input");
769 let bytes = src.as_bytes();
770 let len = src
771 .nrows()
772 .checked_mul(row_bytes)
773 .ok_or_else(|| TrxError::Argument(format!("array '{name}' is too large")))?;
774 dst[byte_cursor..byte_cursor + len].copy_from_slice(bytes);
775 byte_cursor += len;
776 }
777 }
778 Ok(())
779}
780
781fn copy_groups(
782 inputs: &[&AnyTrxFile],
783 group_renames: &[InputGroupRename],
784 specs: &HashMap<String, GroupSpec>,
785 outputs: &mut HashMap<String, DataArray>,
786) -> Result<()> {
787 let mut positions: HashMap<String, usize> =
788 specs.keys().map(|name| (name.clone(), 0)).collect();
789 let mut streamline_base = 0u32;
790 for (input, rename) in inputs.iter().zip(group_renames.iter()) {
791 let src_groups = group_map(input);
792 if src_groups.is_empty() {
793 if let Some(name) = &rename.aggregate_name {
794 let cursor = positions
795 .get_mut(name)
796 .expect("output positions must exist for every group");
797 let dst = outputs
798 .get_mut(name)
799 .expect("output arrays must exist for every group");
800 copy_full_streamline_range(*cursor, input.nb_streamlines(), streamline_base, dst)?;
801 *cursor += input.nb_streamlines();
802 }
803 }
804 for (name, arr) in src_groups {
805 let output_name = rename
806 .renamed_groups
807 .iter()
808 .find_map(|(from, to)| (from == name).then_some(to.as_str()))
809 .unwrap_or(name);
810 let cursor = positions
811 .get_mut(output_name)
812 .expect("output positions must exist for every group");
813 let dst = outputs
814 .get_mut(output_name)
815 .expect("output arrays must exist for every group");
816 let count = arr.nrows();
817 copy_group_with_offset(arr, *cursor, streamline_base, dst)?;
818 *cursor += count;
819 }
820 streamline_base = streamline_base
821 .checked_add(input.nb_streamlines() as u32)
822 .ok_or_else(|| {
823 TrxError::Argument("streamline count overflow during concatenate".into())
824 })?;
825 }
826 Ok(())
827}
828
829fn copy_groups_typed<P: TrxScalar>(
830 inputs: &[&TrxFile<P>],
831 specs: &HashMap<String, GroupSpec>,
832 outputs: &mut HashMap<String, DataArray>,
833) -> Result<()> {
834 let mut positions: HashMap<String, usize> =
835 specs.keys().map(|name| (name.clone(), 0)).collect();
836 let mut streamline_base = 0u32;
837 for input in inputs {
838 for (name, arr) in input.group_arrays() {
839 let cursor = positions
840 .get_mut(name)
841 .expect("output positions must exist for every group");
842 let dst = outputs
843 .get_mut(name)
844 .expect("output arrays must exist for every group");
845 let count = arr.nrows();
846 copy_group_with_offset(arr, *cursor, streamline_base, dst)?;
847 *cursor += count;
848 }
849 streamline_base = streamline_base
850 .checked_add(input.nb_streamlines() as u32)
851 .ok_or_else(|| {
852 TrxError::Argument("streamline count overflow during concatenate".into())
853 })?;
854 }
855 Ok(())
856}
857
858fn copy_full_streamline_range(
859 dst_start: usize,
860 streamline_count: usize,
861 streamline_base: u32,
862 dst: &mut DataArray,
863) -> Result<()> {
864 match dst.dtype() {
865 DType::UInt32 => {
866 let values: &mut [u32] = dst.cast_slice_mut()?;
867 for idx in 0..streamline_count {
868 values[dst_start + idx] =
869 streamline_base.checked_add(idx as u32).ok_or_else(|| {
870 TrxError::Argument("group index overflow during concatenate".into())
871 })?;
872 }
873 Ok(())
874 }
875 other => Err(TrxError::DType(format!(
876 "generated group arrays must use uint32 dtype, got {other}"
877 ))),
878 }
879}
880
881fn copy_group_with_offset(
882 src: &DataArray,
883 dst_start: usize,
884 streamline_base: u32,
885 dst: &mut DataArray,
886) -> Result<()> {
887 match src.dtype() {
888 DType::Int8 => copy_group_typed::<i8>(src.cast_slice(), dst_start, streamline_base, dst),
889 DType::Int16 => copy_group_typed::<i16>(src.cast_slice(), dst_start, streamline_base, dst),
890 DType::Int32 => copy_group_typed::<i32>(src.cast_slice(), dst_start, streamline_base, dst),
891 DType::Int64 => copy_group_typed::<i64>(src.cast_slice(), dst_start, streamline_base, dst),
892 DType::UInt8 => copy_group_typed::<u8>(src.cast_slice(), dst_start, streamline_base, dst),
893 DType::UInt16 => copy_group_typed::<u16>(src.cast_slice(), dst_start, streamline_base, dst),
894 DType::UInt32 => copy_group_typed::<u32>(src.cast_slice(), dst_start, streamline_base, dst),
895 DType::UInt64 => copy_group_typed::<u64>(src.cast_slice(), dst_start, streamline_base, dst),
896 other => Err(TrxError::DType(format!(
897 "group arrays must use integer dtype, got {other}"
898 ))),
899 }
900}
901
902fn copy_group_typed<T>(
903 src: &[T],
904 dst_start: usize,
905 streamline_base: u32,
906 dst: &mut DataArray,
907) -> Result<()>
908where
909 T: Copy + TryFrom<u64> + Into<i128> + bytemuck::Pod,
910 <T as TryFrom<u64>>::Error: std::fmt::Display,
911{
912 let dst_values: &mut [T] = dst.cast_slice_mut()?;
913 let addend = i128::from(streamline_base);
914 for (index, &value) in src.iter().enumerate() {
915 let current = value.into();
916 let shifted = current + addend;
917 let shifted_u64 = u64::try_from(shifted)
918 .map_err(|_| TrxError::Argument("group index underflow during concatenate".into()))?;
919 dst_values[dst_start + index] = T::try_from(shifted_u64).map_err(|err| {
920 TrxError::Argument(format!("group index overflow during concatenate: {err}"))
921 })?;
922 }
923 Ok(())
924}
925
926fn array_map(file: &AnyTrxFile, kind: ArrayKind) -> &HashMap<String, DataArray> {
927 match file {
928 AnyTrxFile::F16(trx) => match kind {
929 ArrayKind::Dps => trx.dps_arrays(),
930 ArrayKind::Dpv => trx.dpv_arrays(),
931 },
932 AnyTrxFile::F32(trx) => match kind {
933 ArrayKind::Dps => trx.dps_arrays(),
934 ArrayKind::Dpv => trx.dpv_arrays(),
935 },
936 AnyTrxFile::F64(trx) => match kind {
937 ArrayKind::Dps => trx.dps_arrays(),
938 ArrayKind::Dpv => trx.dpv_arrays(),
939 },
940 }
941}
942
943fn group_map(file: &AnyTrxFile) -> &HashMap<String, DataArray> {
944 match file {
945 AnyTrxFile::F16(trx) => trx.group_arrays(),
946 AnyTrxFile::F32(trx) => trx.group_arrays(),
947 AnyTrxFile::F64(trx) => trx.group_arrays(),
948 }
949}
950
951fn normalized_group_names(
952 input_count: usize,
953 group_names: &[Option<String>],
954) -> Result<Vec<Option<String>>> {
955 if group_names.is_empty() {
956 return Ok(vec![None; input_count]);
957 }
958 if group_names.len() != input_count {
959 return Err(TrxError::Argument(format!(
960 "input_group_names length {} must match input count {input_count}",
961 group_names.len()
962 )));
963 }
964 Ok(group_names
965 .iter()
966 .map(|value| {
967 value.as_ref().and_then(|name| {
968 let trimmed = name.trim();
969 (!trimmed.is_empty()).then(|| trimmed.to_string())
970 })
971 })
972 .collect())
973}
974
975fn effective_group_rename<I>(group_name: Option<String>, source_groups: I) -> InputGroupRename
976where
977 I: IntoIterator<Item = String>,
978{
979 let source_groups = source_groups.into_iter().collect::<Vec<_>>();
980 if source_groups.is_empty() {
981 return InputGroupRename {
982 aggregate_name: group_name,
983 renamed_groups: Vec::new(),
984 };
985 }
986
987 let renamed_groups = if let Some(group_name) = group_name {
988 source_groups
989 .into_iter()
990 .map(|existing| {
991 let renamed = format!("{group_name}{existing}");
992 (existing, renamed)
993 })
994 .collect()
995 } else {
996 Vec::new()
997 };
998
999 InputGroupRename {
1000 aggregate_name: None,
1001 renamed_groups,
1002 }
1003}
1004
1005#[cfg(test)]
1006mod tests {
1007 use super::*;
1008 use crate::header::Header;
1009 use crate::mmap_backing::vec_to_bytes;
1010
1011 fn sample_header() -> Header {
1012 Header {
1013 voxel_to_rasmm: Header::identity_affine(),
1014 dimensions: [10, 20, 30],
1015 nb_streamlines: 0,
1016 nb_vertices: 0,
1017 extra: Default::default(),
1018 }
1019 }
1020
1021 fn build_trx(
1022 positions: Vec<[f32; 3]>,
1023 offsets: Vec<u32>,
1024 dps: HashMap<String, DataArray>,
1025 dpv: HashMap<String, DataArray>,
1026 groups: HashMap<String, DataArray>,
1027 dpg: HashMap<String, HashMap<String, DataArray>>,
1028 ) -> TrxFile<f32> {
1029 let mut header = sample_header();
1030 header.nb_streamlines = offsets.len().saturating_sub(1) as u64;
1031 header.nb_vertices = positions.len() as u64;
1032 TrxFile::from_parts(TrxParts {
1033 header,
1034 positions_backing: MmapBacking::Owned(vec_to_bytes(positions)),
1035 offsets_backing: MmapBacking::Owned(vec_to_bytes(offsets)),
1036 dps,
1037 dpv,
1038 groups,
1039 dpg,
1040 tempdir: None,
1041 })
1042 }
1043
1044 fn scalar_u32(value: u32) -> DataArray {
1045 DataArray::owned_bytes(vec_to_bytes(vec![value]), 1, DType::UInt32)
1046 }
1047
1048 fn scalar_f32(value: f32) -> DataArray {
1049 DataArray::owned_bytes(vec_to_bytes(vec![value]), 1, DType::Float32)
1050 }
1051
1052 fn vertex_f32(values: Vec<f32>) -> DataArray {
1053 DataArray::owned_bytes(vec_to_bytes(values), 1, DType::Float32)
1054 }
1055
1056 #[test]
1057 fn concatenate_preserves_counts_and_group_union_and_drops_dpg() {
1058 let mut groups_a = HashMap::new();
1059 groups_a.insert(
1060 "left".into(),
1061 DataArray::owned_bytes(vec_to_bytes(vec![0u32]), 1, DType::UInt32),
1062 );
1063 let mut dpg_a = HashMap::new();
1064 dpg_a.insert(
1065 "left".into(),
1066 HashMap::from([(
1067 "color".into(),
1068 DataArray::owned_bytes(vec![1, 2, 3], 3, DType::UInt8),
1069 )]),
1070 );
1071
1072 let a = build_trx(
1073 vec![[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
1074 vec![0, 2],
1075 HashMap::from([("weight".into(), scalar_f32(1.0))]),
1076 HashMap::from([("fa".into(), vertex_f32(vec![0.1, 0.2]))]),
1077 groups_a,
1078 dpg_a,
1079 );
1080
1081 let mut groups_b = HashMap::new();
1082 groups_b.insert(
1083 "right".into(),
1084 DataArray::owned_bytes(vec_to_bytes(vec![0u32]), 1, DType::UInt32),
1085 );
1086 let b = build_trx(
1087 vec![[2.0, 2.0, 2.0]],
1088 vec![0, 1],
1089 HashMap::from([("weight".into(), scalar_f32(2.0))]),
1090 HashMap::from([("fa".into(), vertex_f32(vec![0.3]))]),
1091 groups_b,
1092 HashMap::new(),
1093 );
1094
1095 let any_a = AnyTrxFile::F32(a);
1096 let any_b = AnyTrxFile::F32(b);
1097 let merged =
1098 concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap();
1099
1100 match merged {
1101 AnyTrxFile::F32(trx) => {
1102 assert_eq!(trx.nb_streamlines(), 2);
1103 assert_eq!(trx.nb_vertices(), 3);
1104 assert_eq!(trx.header().dimensions, [10, 20, 30]);
1105 assert!(trx.is_file_backed());
1106 assert_eq!(trx.group("left").unwrap(), &[0]);
1107 assert_eq!(trx.group("right").unwrap(), &[1]);
1108 assert!(trx.dpg_group_names().is_empty());
1109 }
1110 _ => panic!("expected float32 output"),
1111 }
1112 }
1113
1114 #[test]
1115 fn concatenate_errors_on_missing_dps_without_delete_flag() {
1116 let a = build_trx(
1117 vec![[0.0, 0.0, 0.0]],
1118 vec![0, 1],
1119 HashMap::from([("weight".into(), scalar_f32(1.0))]),
1120 HashMap::new(),
1121 HashMap::new(),
1122 HashMap::new(),
1123 );
1124 let b = build_trx(
1125 vec![[1.0, 1.0, 1.0]],
1126 vec![0, 1],
1127 HashMap::new(),
1128 HashMap::new(),
1129 HashMap::new(),
1130 HashMap::new(),
1131 );
1132 let any_a = AnyTrxFile::F32(a);
1133 let any_b = AnyTrxFile::F32(b);
1134 let err =
1135 concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap_err();
1136 assert!(err.to_string().contains("dps key 'weight'"));
1137 }
1138
1139 #[test]
1140 fn concatenate_delete_dps_drops_category() {
1141 let a = build_trx(
1142 vec![[0.0, 0.0, 0.0]],
1143 vec![0, 1],
1144 HashMap::from([("weight".into(), scalar_f32(1.0))]),
1145 HashMap::new(),
1146 HashMap::new(),
1147 HashMap::new(),
1148 );
1149 let b = build_trx(
1150 vec![[1.0, 1.0, 1.0]],
1151 vec![0, 1],
1152 HashMap::new(),
1153 HashMap::new(),
1154 HashMap::new(),
1155 HashMap::new(),
1156 );
1157 let any_a = AnyTrxFile::F32(a);
1158 let any_b = AnyTrxFile::F32(b);
1159 let merged = concatenate_any_trx(
1160 &[&any_a, &any_b],
1161 &ConcatenateOptions {
1162 delete_dps: true,
1163 ..Default::default()
1164 },
1165 )
1166 .unwrap();
1167 match merged {
1168 AnyTrxFile::F32(trx) => assert!(trx.dps_names().is_empty()),
1169 _ => panic!("expected float32 output"),
1170 }
1171 }
1172
1173 #[test]
1174 fn concatenate_errors_on_dpv_column_mismatch() {
1175 let a = build_trx(
1176 vec![[0.0, 0.0, 0.0]],
1177 vec![0, 1],
1178 HashMap::new(),
1179 HashMap::from([(
1180 "fa".into(),
1181 DataArray::owned_bytes(vec_to_bytes(vec![0.1f32]), 1, DType::Float32),
1182 )]),
1183 HashMap::new(),
1184 HashMap::new(),
1185 );
1186 let b = build_trx(
1187 vec![[1.0, 1.0, 1.0]],
1188 vec![0, 1],
1189 HashMap::new(),
1190 HashMap::from([(
1191 "fa".into(),
1192 DataArray::owned_bytes(vec_to_bytes(vec![0.2f32, 0.3f32]), 2, DType::Float32),
1193 )]),
1194 HashMap::new(),
1195 HashMap::new(),
1196 );
1197 let any_a = AnyTrxFile::F32(a);
1198 let any_b = AnyTrxFile::F32(b);
1199 let err =
1200 concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap_err();
1201 assert!(err.to_string().contains("column counts"));
1202 }
1203
1204 #[test]
1205 fn concatenate_errors_on_group_dtype_mismatch() {
1206 let a = build_trx(
1207 vec![[0.0, 0.0, 0.0]],
1208 vec![0, 1],
1209 HashMap::new(),
1210 HashMap::new(),
1211 HashMap::from([("bundle".into(), scalar_u32(0))]),
1212 HashMap::new(),
1213 );
1214 let b = build_trx(
1215 vec![[1.0, 1.0, 1.0]],
1216 vec![0, 1],
1217 HashMap::new(),
1218 HashMap::new(),
1219 HashMap::from([(
1220 "bundle".into(),
1221 DataArray::owned_bytes(vec_to_bytes(vec![0u16]), 1, DType::UInt16),
1222 )]),
1223 HashMap::new(),
1224 );
1225 let any_a = AnyTrxFile::F32(a);
1226 let any_b = AnyTrxFile::F32(b);
1227 let err =
1228 concatenate_any_trx(&[&any_a, &any_b], &ConcatenateOptions::default()).unwrap_err();
1229 assert!(err.to_string().contains("group key 'bundle'"));
1230 }
1231
1232 #[test]
1233 fn concatenate_respects_positions_dtype_override() {
1234 let a = build_trx(
1235 vec![[0.0, 0.0, 0.0]],
1236 vec![0, 1],
1237 HashMap::new(),
1238 HashMap::new(),
1239 HashMap::new(),
1240 HashMap::new(),
1241 );
1242 let b = build_trx(
1243 vec![[1.0, 1.0, 1.0]],
1244 vec![0, 1],
1245 HashMap::new(),
1246 HashMap::new(),
1247 HashMap::new(),
1248 HashMap::new(),
1249 );
1250 let any_a = AnyTrxFile::F32(a);
1251 let any_b = AnyTrxFile::F32(b);
1252 let merged = concatenate_any_trx(
1253 &[&any_a, &any_b],
1254 &ConcatenateOptions {
1255 positions_dtype: Some(DType::Float16),
1256 ..Default::default()
1257 },
1258 )
1259 .unwrap();
1260 assert!(matches!(merged, AnyTrxFile::F16(_)));
1261 }
1262
1263 #[test]
1264 fn concatenate_errors_on_group_name_length_mismatch() {
1265 let a = build_trx(
1266 vec![[0.0, 0.0, 0.0]],
1267 vec![0, 1],
1268 HashMap::new(),
1269 HashMap::new(),
1270 HashMap::new(),
1271 HashMap::new(),
1272 );
1273 let b = build_trx(
1274 vec![[1.0, 1.0, 1.0]],
1275 vec![0, 1],
1276 HashMap::new(),
1277 HashMap::new(),
1278 HashMap::new(),
1279 HashMap::new(),
1280 );
1281 let any_a = AnyTrxFile::F32(a);
1282 let any_b = AnyTrxFile::F32(b);
1283 let err = concatenate_any_trx(
1284 &[&any_a, &any_b],
1285 &ConcatenateOptions {
1286 input_group_names: vec![Some("bundle".into())],
1287 ..Default::default()
1288 },
1289 )
1290 .unwrap_err();
1291 assert!(err.to_string().contains("input_group_names length"));
1292 }
1293
1294 #[test]
1295 fn concatenate_group_name_creates_aggregate_group_when_input_has_no_groups() {
1296 let a = build_trx(
1297 vec![[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]],
1298 vec![0, 1, 2],
1299 HashMap::new(),
1300 HashMap::new(),
1301 HashMap::new(),
1302 HashMap::new(),
1303 );
1304 let b = build_trx(
1305 vec![[2.0, 0.0, 0.0]],
1306 vec![0, 1],
1307 HashMap::new(),
1308 HashMap::new(),
1309 HashMap::new(),
1310 HashMap::new(),
1311 );
1312 let any_a = AnyTrxFile::F32(a);
1313 let any_b = AnyTrxFile::F32(b);
1314 let merged = concatenate_any_trx(
1315 &[&any_a, &any_b],
1316 &ConcatenateOptions {
1317 input_group_names: vec![Some("first".into()), None],
1318 ..Default::default()
1319 },
1320 )
1321 .unwrap();
1322
1323 match merged {
1324 AnyTrxFile::F32(trx) => {
1325 assert_eq!(trx.group("first").unwrap(), &[0, 1]);
1326 }
1327 _ => panic!("expected float32 output"),
1328 }
1329 }
1330
1331 #[test]
1332 fn concatenate_group_name_prefixes_existing_groups_without_aggregate() {
1333 let a = build_trx(
1334 vec![[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
1335 vec![0, 1, 2],
1336 HashMap::new(),
1337 HashMap::new(),
1338 HashMap::from([(
1339 "bundle".into(),
1340 DataArray::owned_bytes(vec_to_bytes(vec![0u32, 1u32]), 1, DType::UInt32),
1341 )]),
1342 HashMap::new(),
1343 );
1344 let b = build_trx(
1345 vec![[2.0, 2.0, 2.0]],
1346 vec![0, 1],
1347 HashMap::new(),
1348 HashMap::new(),
1349 HashMap::new(),
1350 HashMap::new(),
1351 );
1352 let any_a = AnyTrxFile::F32(a);
1353 let any_b = AnyTrxFile::F32(b);
1354 let merged = concatenate_any_trx(
1355 &[&any_a, &any_b],
1356 &ConcatenateOptions {
1357 input_group_names: vec![Some("Prefix".into()), None],
1358 ..Default::default()
1359 },
1360 )
1361 .unwrap();
1362
1363 match merged {
1364 AnyTrxFile::F32(trx) => {
1365 assert_eq!(trx.group("Prefixbundle").unwrap(), &[0, 1]);
1366 assert!(!trx.group_names().contains(&"Prefix"));
1367 }
1368 _ => panic!("expected float32 output"),
1369 }
1370 }
1371
1372 #[test]
1373 fn concatenate_group_name_blank_entry_is_ignored() {
1374 let a = build_trx(
1375 vec![[0.0, 0.0, 0.0]],
1376 vec![0, 1],
1377 HashMap::new(),
1378 HashMap::new(),
1379 HashMap::new(),
1380 HashMap::new(),
1381 );
1382 let b = build_trx(
1383 vec![[1.0, 1.0, 1.0]],
1384 vec![0, 1],
1385 HashMap::new(),
1386 HashMap::new(),
1387 HashMap::new(),
1388 HashMap::new(),
1389 );
1390 let any_a = AnyTrxFile::F32(a);
1391 let any_b = AnyTrxFile::F32(b);
1392 let merged = concatenate_any_trx(
1393 &[&any_a, &any_b],
1394 &ConcatenateOptions {
1395 input_group_names: vec![Some(" ".into()), None],
1396 ..Default::default()
1397 },
1398 )
1399 .unwrap();
1400
1401 match merged {
1402 AnyTrxFile::F32(trx) => assert!(trx.group_names().is_empty()),
1403 _ => panic!("expected float32 output"),
1404 }
1405 }
1406
1407 #[test]
1408 fn concatenate_overlapping_prefixed_group_names_union_members() {
1409 let a = build_trx(
1410 vec![[0.0, 0.0, 0.0]],
1411 vec![0, 1],
1412 HashMap::new(),
1413 HashMap::new(),
1414 HashMap::from([("One".into(), scalar_u32(0))]),
1415 HashMap::new(),
1416 );
1417 let b = build_trx(
1418 vec![[1.0, 1.0, 1.0]],
1419 vec![0, 1],
1420 HashMap::new(),
1421 HashMap::new(),
1422 HashMap::from([("One".into(), scalar_u32(0))]),
1423 HashMap::new(),
1424 );
1425 let any_a = AnyTrxFile::F32(a);
1426 let any_b = AnyTrxFile::F32(b);
1427 let merged = concatenate_any_trx(
1428 &[&any_a, &any_b],
1429 &ConcatenateOptions {
1430 input_group_names: vec![Some("Same".into()), Some("Same".into())],
1431 ..Default::default()
1432 },
1433 )
1434 .unwrap();
1435
1436 match merged {
1437 AnyTrxFile::F32(trx) => assert_eq!(trx.group("SameOne").unwrap(), &[0, 1]),
1438 _ => panic!("expected float32 output"),
1439 }
1440 }
1441}