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