gemo-tractography 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 Amin Barati
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,212 @@
1
+ Metadata-Version: 2.4
2
+ Name: gemo-tractography
3
+ Version: 0.1.0
4
+ Summary: Geometric & morphological deep learning classifier for tractography streamlines
5
+ Author: Amin Barati
6
+ License: MIT
7
+ Project-URL: Homepage, https://github.com/amin-barati/GEMO
8
+ Project-URL: Repository, https://github.com/amin-barati/GEMO
9
+ Keywords: tractography,diffusion-mri,deep-learning,neuroimaging,streamlines
10
+ Classifier: Programming Language :: Python :: 3
11
+ Classifier: License :: OSI Approved :: MIT License
12
+ Classifier: Operating System :: OS Independent
13
+ Classifier: Intended Audience :: Science/Research
14
+ Classifier: Topic :: Scientific/Engineering :: Medical Science Apps.
15
+ Requires-Python: >=3.9
16
+ Description-Content-Type: text/markdown
17
+ License-File: LICENSE
18
+ Requires-Dist: torch>=2.0
19
+ Requires-Dist: numpy
20
+ Requires-Dist: nibabel
21
+ Requires-Dist: h5py
22
+ Requires-Dist: tqdm
23
+ Requires-Dist: scipy
24
+ Requires-Dist: scikit-image
25
+ Dynamic: license-file
26
+
27
+ # **GEMO: A deep learning method for brain fiber classification and tract segmentation using geometrical and morphological features**
28
+
29
+ PyTorch implementation of GEMO, a deep learning framework for white matter streamline classification and tract segmentation using convolutional neural networks together with handcrafted geometric and morphological features.
30
+
31
+ **Overview**
32
+
33
+ Diffusion-weighted magnetic resonance imaging (dMRI) is often used to study brain structure. One of the important applications made possible by dMRI is streamline tractography that is utilized for structural connectivity evaluation and neurosurgical planning. GEMO is a supervised deep learning framework for white matter streamline classification that combines convolutional neural network (CNN) features with handcrafted geometric and morphological descriptors to improve classification performance. This approach benefits from GEometrical and MOrphological features in addition to the features extracted from a convolutional neural network for improving the final classification performance. Streamlines are transformed from three-dimensional coordinate space into two-dimensional color-encoded images using the “xyz2RGB” mapping method and are subsequently fed into a convolutional neural network for learning and classification. This method performs the streamline classification from a whole brain tractogram, only focusing on each streamline’s features, which include those provided from CNN in addition to geometric and morphologic features.
34
+
35
+ ![alt text](Figures/Graphical_abstract.png)
36
+
37
+ **Complete list of output classes in GEMO**
38
+
39
+ - AF_L — Left Arcuate Fasciculus
40
+ - AF_R — Right Arcuate Fasciculus
41
+ - CC_Fr_1 — Corpus Callosum, Frontal Region 1
42
+ - CC_Fr_2 — Corpus Callosum, Frontal Region 2
43
+ - CC_Oc — Corpus Callosum, Occipital Region
44
+ - CC_Pa — Corpus Callosum, Parietal Region
45
+ - CC_Pr_Po — Corpus Callosum, Precentral/Postcentral Region
46
+ - CG_L — Left Cingulum
47
+ - CG_R — Right Cingulum
48
+ - FAT_L — Left Frontal Aslant Tract
49
+ - FAT_R — Right Frontal Aslant Tract
50
+ - FPT_L — Left Frontopontine Tract
51
+ - FPT_R — Right Frontopontine Tract
52
+ - IFOF_L — Left Inferior Fronto-Occipital Fasciculus
53
+ - IFOF_R — Right Inferior Fronto-Occipital Fasciculus
54
+ - ILF_L — Left Inferior Longitudinal Fasciculus
55
+ - ILF_R — Right Inferior Longitudinal Fasciculus
56
+ - MCP — Middle Cerebellar Peduncle
57
+ - MdLF_L — Left Middle Longitudinal Fasciculus
58
+ - MdLF_R — Right Middle Longitudinal Fasciculus
59
+ - POPT_L — Left Parieto-Occipital Pontine Tract
60
+ - POPT_R — Right Parieto-Occipital Pontine Tract
61
+ - PYT_L — Left Pyramidal (Corticospinal) Tract
62
+ - PYT_R — Right Pyramidal (Corticospinal) Tract
63
+ - SLF_L — Left Superior Longitudinal Fasciculus
64
+ - SLF_R — Right Superior Longitudinal Fasciculus
65
+ - UF_L — Left Uncinate Fasciculus
66
+ - UF_R — Right Uncinate Fasciculus
67
+ - OR_ML_L — Left Optic Radiation (Meyer's Loop)
68
+ - OR_ML_R — Right Optic Radiation (Meyer's Loop)
69
+
70
+
71
+
72
+ # Installation
73
+
74
+ Install directly from GitHub:
75
+
76
+ ```bash
77
+ pip install git+https://github.com/amin-barati/GEMO.git
78
+ ```
79
+
80
+ This requires `git` to be available on your machine.
81
+
82
+ # Usage: classify a tractogram
83
+
84
+ ```bash
85
+ gemo-infer --trk-path wholebrain.trk --output-dir classified_output
86
+ ```
87
+
88
+ By default this uses the checkpoint and label map bundled with the package
89
+ (`gemo/checkpoints/best.pt`, `gemo/checkpoints/label_map.json`). To use a
90
+ different, custom-trained model instead:
91
+
92
+ ```bash
93
+ gemo-infer --trk-path wholebrain.trk \
94
+ --checkpoint /path/to/best.pt \
95
+ --label-map /path/to/label_map.json \
96
+ --output-dir classified_output \
97
+ --threshold 0.70
98
+ ```
99
+
100
+ This writes one `.trk` file per predicted tract class, plus one
101
+ `_Unknown.trk` file for streamlines whose top softmax probability falls
102
+ below `--threshold`.
103
+
104
+
105
+ # xyz2RGB
106
+
107
+ The `xyz2RGB` module converts each streamline of a tractogram into a RGB image by mapping the normalized x, y, and z coordinates to the red, green, and blue color channels, respectively. These images provide a compact representation of streamline geometry and serve as the input to the convolutional neural network (CNN) used in GEMO. The following example converts all streamlines contained in a single `.trk` file into RGB images.
108
+
109
+ ![image_xyz2RGB](Figures/xyz2RGB.png)
110
+
111
+
112
+ **Generate RGB images from a single tractogram**
113
+
114
+ An example of how `xyz2RGB` maps each streamline to an RGB image is provided in the following code.
115
+
116
+ ```python
117
+ from xyz2RGB_test import trk_to_images
118
+
119
+ images = trk_to_images("Sample_Tract.trk", output_dir="output_images")
120
+ ```
121
+
122
+ **TCK to TRK Conversion**
123
+
124
+ If your tractogram files are in `.tck` format, convert them to `.trk` format before starting training or inference. The `tck_to_trk.py` utility supports both individual files and entire directories.
125
+
126
+ Convert a single `.tck` file:
127
+ ```bash
128
+ python tck_to_trk.py --input AF_left.tck --output AF_left.trk
129
+ ```
130
+
131
+ Convert all .tck files in a directory:
132
+
133
+ ```bash
134
+ python tck_to_trk.py --input "TCK_directory" --output "TRK_directory"
135
+ ```
136
+
137
+ # Training on your own data
138
+
139
+ After generating the RGB images and extracting the handcrafted features, the model can be trained using the provided training script. During training, the RGB images are processed by a CNN, the handcrafted features are encoded using a multilayer perceptron (MLP), and the learned representations are fused to classify each streamline into its corresponding white matter tract. The command below starts the complete training pipeline.
140
+
141
+
142
+ # Requirements
143
+ python==3.12
144
+ numpy==2.2.6
145
+ scipy==1.16.0
146
+ torch==2.8.0
147
+ torchvision==0.23.0
148
+ nibabel==5.3.2
149
+ h5py==3.14.0
150
+ tqdm==4.67.1
151
+ scikit-learn==1.7.1
152
+ matplotlib==3.10.5
153
+
154
+ ```bash
155
+ pip install -r requirements.txt
156
+ ```
157
+
158
+ GEMO/
159
+ ├── train.py
160
+ ├── streamline_model.py
161
+ ├── dataset.py
162
+ ├── streamline_features.py
163
+ ├── xyz2RGB.py
164
+ ├── utils.py
165
+ ├── config.py
166
+ ├── inference.py
167
+ ├── requirements.txt
168
+ ├── TRK_directory/
169
+ │ ├── Tract_01.trk
170
+ │ ├── ...
171
+ │ └── Tract_999.trk
172
+ ├── bounds.h5
173
+ └── features.h5
174
+
175
+ ## 1. Precompute subject bounds
176
+
177
+ ```python
178
+ from streamline_features import compute_and_save_bounds_metadata
179
+ compute_and_save_bounds_metadata("TRK_directory", "bounds.h5")
180
+ ```
181
+
182
+
183
+
184
+ ## 2. Extract handcrafted features
185
+
186
+ **Extract geometric and morphological features**
187
+
188
+ GEMO utilizes several handcrafted geometric and morphological descriptors for each streamline, including length, curvature, tortuosity, spectral entropy, fractal dimension, and lacunarity. These features are computed once and stored in an HDF5 (`.h5`) file, allowing efficient loading during training without repeated feature computation.
189
+
190
+ ```python
191
+ from streamline_features import extract_features_from_directory
192
+ extract_features_from_directory("TRK_directory", "features.h5", bounds_h5_path="bounds.h5")
193
+ ```
194
+ ## 3. Train model
195
+
196
+ ```bash
197
+ python train.py --trk-dir TRK_directory --features-h5 features.h5 --bounds-h5 bounds.h5
198
+ ```
199
+ ## 4. Inference
200
+
201
+ ```bash
202
+ python inference.py --trk-path wholebrain.trk --checkpoint checkpoints/best.pt --label-map checkpoints/label_map.json --output-dir classified_output --threshold 0.70
203
+ ```
204
+
205
+ # Citation
206
+ GEMO is the code for the following paper; if you use this repository in your research, please cite:
207
+
208
+ GEMO: A deep learning method for brain fiber classification and tract segmentation using geometrical and morphological features
209
+
210
+ Journal: Academic Radiology
211
+
212
+ DOI:https://doi.org/10.1016/j.acra.2026.07.024
@@ -0,0 +1,186 @@
1
+ # **GEMO: A deep learning method for brain fiber classification and tract segmentation using geometrical and morphological features**
2
+
3
+ PyTorch implementation of GEMO, a deep learning framework for white matter streamline classification and tract segmentation using convolutional neural networks together with handcrafted geometric and morphological features.
4
+
5
+ **Overview**
6
+
7
+ Diffusion-weighted magnetic resonance imaging (dMRI) is often used to study brain structure. One of the important applications made possible by dMRI is streamline tractography that is utilized for structural connectivity evaluation and neurosurgical planning. GEMO is a supervised deep learning framework for white matter streamline classification that combines convolutional neural network (CNN) features with handcrafted geometric and morphological descriptors to improve classification performance. This approach benefits from GEometrical and MOrphological features in addition to the features extracted from a convolutional neural network for improving the final classification performance. Streamlines are transformed from three-dimensional coordinate space into two-dimensional color-encoded images using the “xyz2RGB” mapping method and are subsequently fed into a convolutional neural network for learning and classification. This method performs the streamline classification from a whole brain tractogram, only focusing on each streamline’s features, which include those provided from CNN in addition to geometric and morphologic features.
8
+
9
+ ![alt text](Figures/Graphical_abstract.png)
10
+
11
+ **Complete list of output classes in GEMO**
12
+
13
+ - AF_L — Left Arcuate Fasciculus
14
+ - AF_R — Right Arcuate Fasciculus
15
+ - CC_Fr_1 — Corpus Callosum, Frontal Region 1
16
+ - CC_Fr_2 — Corpus Callosum, Frontal Region 2
17
+ - CC_Oc — Corpus Callosum, Occipital Region
18
+ - CC_Pa — Corpus Callosum, Parietal Region
19
+ - CC_Pr_Po — Corpus Callosum, Precentral/Postcentral Region
20
+ - CG_L — Left Cingulum
21
+ - CG_R — Right Cingulum
22
+ - FAT_L — Left Frontal Aslant Tract
23
+ - FAT_R — Right Frontal Aslant Tract
24
+ - FPT_L — Left Frontopontine Tract
25
+ - FPT_R — Right Frontopontine Tract
26
+ - IFOF_L — Left Inferior Fronto-Occipital Fasciculus
27
+ - IFOF_R — Right Inferior Fronto-Occipital Fasciculus
28
+ - ILF_L — Left Inferior Longitudinal Fasciculus
29
+ - ILF_R — Right Inferior Longitudinal Fasciculus
30
+ - MCP — Middle Cerebellar Peduncle
31
+ - MdLF_L — Left Middle Longitudinal Fasciculus
32
+ - MdLF_R — Right Middle Longitudinal Fasciculus
33
+ - POPT_L — Left Parieto-Occipital Pontine Tract
34
+ - POPT_R — Right Parieto-Occipital Pontine Tract
35
+ - PYT_L — Left Pyramidal (Corticospinal) Tract
36
+ - PYT_R — Right Pyramidal (Corticospinal) Tract
37
+ - SLF_L — Left Superior Longitudinal Fasciculus
38
+ - SLF_R — Right Superior Longitudinal Fasciculus
39
+ - UF_L — Left Uncinate Fasciculus
40
+ - UF_R — Right Uncinate Fasciculus
41
+ - OR_ML_L — Left Optic Radiation (Meyer's Loop)
42
+ - OR_ML_R — Right Optic Radiation (Meyer's Loop)
43
+
44
+
45
+
46
+ # Installation
47
+
48
+ Install directly from GitHub:
49
+
50
+ ```bash
51
+ pip install git+https://github.com/amin-barati/GEMO.git
52
+ ```
53
+
54
+ This requires `git` to be available on your machine.
55
+
56
+ # Usage: classify a tractogram
57
+
58
+ ```bash
59
+ gemo-infer --trk-path wholebrain.trk --output-dir classified_output
60
+ ```
61
+
62
+ By default this uses the checkpoint and label map bundled with the package
63
+ (`gemo/checkpoints/best.pt`, `gemo/checkpoints/label_map.json`). To use a
64
+ different, custom-trained model instead:
65
+
66
+ ```bash
67
+ gemo-infer --trk-path wholebrain.trk \
68
+ --checkpoint /path/to/best.pt \
69
+ --label-map /path/to/label_map.json \
70
+ --output-dir classified_output \
71
+ --threshold 0.70
72
+ ```
73
+
74
+ This writes one `.trk` file per predicted tract class, plus one
75
+ `_Unknown.trk` file for streamlines whose top softmax probability falls
76
+ below `--threshold`.
77
+
78
+
79
+ # xyz2RGB
80
+
81
+ The `xyz2RGB` module converts each streamline of a tractogram into a RGB image by mapping the normalized x, y, and z coordinates to the red, green, and blue color channels, respectively. These images provide a compact representation of streamline geometry and serve as the input to the convolutional neural network (CNN) used in GEMO. The following example converts all streamlines contained in a single `.trk` file into RGB images.
82
+
83
+ ![image_xyz2RGB](Figures/xyz2RGB.png)
84
+
85
+
86
+ **Generate RGB images from a single tractogram**
87
+
88
+ An example of how `xyz2RGB` maps each streamline to an RGB image is provided in the following code.
89
+
90
+ ```python
91
+ from xyz2RGB_test import trk_to_images
92
+
93
+ images = trk_to_images("Sample_Tract.trk", output_dir="output_images")
94
+ ```
95
+
96
+ **TCK to TRK Conversion**
97
+
98
+ If your tractogram files are in `.tck` format, convert them to `.trk` format before starting training or inference. The `tck_to_trk.py` utility supports both individual files and entire directories.
99
+
100
+ Convert a single `.tck` file:
101
+ ```bash
102
+ python tck_to_trk.py --input AF_left.tck --output AF_left.trk
103
+ ```
104
+
105
+ Convert all .tck files in a directory:
106
+
107
+ ```bash
108
+ python tck_to_trk.py --input "TCK_directory" --output "TRK_directory"
109
+ ```
110
+
111
+ # Training on your own data
112
+
113
+ After generating the RGB images and extracting the handcrafted features, the model can be trained using the provided training script. During training, the RGB images are processed by a CNN, the handcrafted features are encoded using a multilayer perceptron (MLP), and the learned representations are fused to classify each streamline into its corresponding white matter tract. The command below starts the complete training pipeline.
114
+
115
+
116
+ # Requirements
117
+ python==3.12
118
+ numpy==2.2.6
119
+ scipy==1.16.0
120
+ torch==2.8.0
121
+ torchvision==0.23.0
122
+ nibabel==5.3.2
123
+ h5py==3.14.0
124
+ tqdm==4.67.1
125
+ scikit-learn==1.7.1
126
+ matplotlib==3.10.5
127
+
128
+ ```bash
129
+ pip install -r requirements.txt
130
+ ```
131
+
132
+ GEMO/
133
+ ├── train.py
134
+ ├── streamline_model.py
135
+ ├── dataset.py
136
+ ├── streamline_features.py
137
+ ├── xyz2RGB.py
138
+ ├── utils.py
139
+ ├── config.py
140
+ ├── inference.py
141
+ ├── requirements.txt
142
+ ├── TRK_directory/
143
+ │ ├── Tract_01.trk
144
+ │ ├── ...
145
+ │ └── Tract_999.trk
146
+ ├── bounds.h5
147
+ └── features.h5
148
+
149
+ ## 1. Precompute subject bounds
150
+
151
+ ```python
152
+ from streamline_features import compute_and_save_bounds_metadata
153
+ compute_and_save_bounds_metadata("TRK_directory", "bounds.h5")
154
+ ```
155
+
156
+
157
+
158
+ ## 2. Extract handcrafted features
159
+
160
+ **Extract geometric and morphological features**
161
+
162
+ GEMO utilizes several handcrafted geometric and morphological descriptors for each streamline, including length, curvature, tortuosity, spectral entropy, fractal dimension, and lacunarity. These features are computed once and stored in an HDF5 (`.h5`) file, allowing efficient loading during training without repeated feature computation.
163
+
164
+ ```python
165
+ from streamline_features import extract_features_from_directory
166
+ extract_features_from_directory("TRK_directory", "features.h5", bounds_h5_path="bounds.h5")
167
+ ```
168
+ ## 3. Train model
169
+
170
+ ```bash
171
+ python train.py --trk-dir TRK_directory --features-h5 features.h5 --bounds-h5 bounds.h5
172
+ ```
173
+ ## 4. Inference
174
+
175
+ ```bash
176
+ python inference.py --trk-path wholebrain.trk --checkpoint checkpoints/best.pt --label-map checkpoints/label_map.json --output-dir classified_output --threshold 0.70
177
+ ```
178
+
179
+ # Citation
180
+ GEMO is the code for the following paper; if you use this repository in your research, please cite:
181
+
182
+ GEMO: A deep learning method for brain fiber classification and tract segmentation using geometrical and morphological features
183
+
184
+ Journal: Academic Radiology
185
+
186
+ DOI:https://doi.org/10.1016/j.acra.2026.07.024
@@ -0,0 +1,69 @@
1
+ """
2
+ ==============================================================================
3
+ GEMO: A Deep Learning Method for Brain Fiber Classification and Tract
4
+ Segmentation Using Geometrical and Morphological Features
5
+
6
+ DOI: https://doi.org/10.1016/j.acra.2026.07.024
7
+
8
+ Email : Amin_br@yahoo.com
9
+ GitHub : https://github.com/amin-barati/GEMO
10
+
11
+ ==============================================================================
12
+
13
+
14
+ Command-line tools
15
+ ------------------------------------------------------
16
+ gemo-train Train a StreamlineClassifier on a directory of .trk files.
17
+ gemo-infer Classify a whole-brain .trk file with a trained checkpoint
18
+
19
+
20
+ Python API
21
+ ----------
22
+ >>> from gemo import StreamlineClassifier, classify_tractogram
23
+ >>> model = StreamlineClassifier(num_handcrafted_features=6, num_classes=30)
24
+ """
25
+
26
+ from .streamline_model import (
27
+ StreamlineClassifier,
28
+ CNNFeatureExtractor,
29
+ HandcraftedFeatureEncoder,
30
+ FeatureFusion,
31
+ ClassificationHead,
32
+ StreamlineDataset,
33
+ build_label_mapping,
34
+ classify_tractogram,
35
+ count_trainable_parameters,
36
+ )
37
+ from .xyz2RGB import (
38
+ streamline_to_rgb,
39
+ load_bounds_metadata_h5,
40
+ save_bounds_metadata_h5,
41
+ group_trk_files_by_subject,
42
+ )
43
+ from .streamline_features import (
44
+ extract_features_from_directory,
45
+ compute_and_save_bounds_metadata,
46
+ )
47
+ from .config import Config
48
+
49
+ __version__ = "0.1.0"
50
+
51
+ __all__ = [
52
+ "StreamlineClassifier",
53
+ "CNNFeatureExtractor",
54
+ "HandcraftedFeatureEncoder",
55
+ "FeatureFusion",
56
+ "ClassificationHead",
57
+ "StreamlineDataset",
58
+ "build_label_mapping",
59
+ "classify_tractogram",
60
+ "count_trainable_parameters",
61
+ "streamline_to_rgb",
62
+ "load_bounds_metadata_h5",
63
+ "save_bounds_metadata_h5",
64
+ "group_trk_files_by_subject",
65
+ "extract_features_from_directory",
66
+ "compute_and_save_bounds_metadata",
67
+ "Config",
68
+ "__version__",
69
+ ]
@@ -0,0 +1,32 @@
1
+ {
2
+ "AF_L": 0,
3
+ "AF_R": 1,
4
+ "CC_Fr_1": 2,
5
+ "CC_Fr_2": 3,
6
+ "CC_Oc": 4,
7
+ "CC_Pa": 5,
8
+ "CC_Pr_Po": 6,
9
+ "CG_L": 7,
10
+ "CG_R": 8,
11
+ "FAT_L": 9,
12
+ "FAT_R": 10,
13
+ "FPT_L": 11,
14
+ "FPT_R": 12,
15
+ "IFOF_L": 13,
16
+ "IFOF_R": 14,
17
+ "ILF_L": 15,
18
+ "ILF_R": 16,
19
+ "MCP": 17,
20
+ "MdLF_L": 18,
21
+ "MdLF_R": 19,
22
+ "OR_ML_L": 20,
23
+ "OR_ML_R": 21,
24
+ "POPT_L": 22,
25
+ "POPT_R": 23,
26
+ "PYT_L": 24,
27
+ "PYT_R": 25,
28
+ "SLF_L": 26,
29
+ "SLF_R": 27,
30
+ "UF_L": 28,
31
+ "UF_R": 29
32
+ }
@@ -0,0 +1,87 @@
1
+ """
2
+ config
3
+ ======
4
+
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from dataclasses import dataclass, field
10
+ from typing import Optional, Tuple
11
+
12
+
13
+ @dataclass
14
+ class DataConfig:
15
+ trk_dir: str = "trk_dir"
16
+ features_h5_path: str = "features.h5"
17
+ bounds_h5_path: str = "bounds.h5"
18
+ separator: str = "__"
19
+ index_width: int = 7
20
+ grid_size: int = 64
21
+ scales: Tuple[int, ...] = (1, 2, 4, 8, 16)
22
+ feature_names: Tuple[str, ...] = (
23
+ "length",
24
+ "curvature",
25
+ "tortuosity",
26
+ "spectral_entropy",
27
+ "fractal_dimension",
28
+ "lacunarity",
29
+ )
30
+ truncate: bool = True
31
+
32
+ shuffle_buffer_size: int = 2000
33
+
34
+ max_streamlines_per_class_per_epoch: Optional[int] = 100_000
35
+ val_fraction: float = 0.2
36
+ split_seed: int = 42
37
+
38
+
39
+ @dataclass
40
+ class ModelConfig:
41
+ num_handcrafted_features: int = 6
42
+ num_classes: int = 30
43
+ cnn_base_channels: int = 32
44
+ cnn_embedding_dim: int = 128
45
+ feature_hidden_dims: Tuple[int, ...] = (32, 64)
46
+ feature_embedding_dim: int = 64
47
+ classifier_hidden_dims: Tuple[int, ...] = (256, 128)
48
+ dropout: float = 0.2
49
+ fusion_mode: str = "concat"
50
+
51
+
52
+ @dataclass
53
+ class TrainConfig:
54
+ batch_size: int = 256
55
+ num_workers: int = 4
56
+ num_epochs: int = 100
57
+ learning_rate: float = 1e-4
58
+ weight_decay: float = 5e-4
59
+ lr_step_size: int = 10
60
+ lr_gamma: float = 0.5
61
+ grad_clip_norm: float = 5.0
62
+ device: str = "auto"
63
+ checkpoint_dir: str = "checkpoints"
64
+ log_every_n_steps: int = 50
65
+ val_every_n_epochs: int = 1
66
+
67
+ class_weight_scheme: str = "inverse_sqrt"
68
+ seed: int = 0
69
+
70
+
71
+ @dataclass
72
+ class InferenceConfig:
73
+ threshold: float = 0.7
74
+ batch_size: int = 256
75
+ device: str = "auto"
76
+ output_dir: str = "classified_output"
77
+
78
+
79
+ @dataclass
80
+ class Config:
81
+ data: DataConfig = field(default_factory=DataConfig)
82
+ model: ModelConfig = field(default_factory=ModelConfig)
83
+ train: TrainConfig = field(default_factory=TrainConfig)
84
+ inference: InferenceConfig = field(default_factory=InferenceConfig)
85
+
86
+
87
+ __all__ = ["DataConfig", "ModelConfig", "TrainConfig", "InferenceConfig", "Config"]