marn 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.
Files changed (68) hide show
  1. marn-0.1.0/LICENSE +161 -0
  2. marn-0.1.0/PKG-INFO +125 -0
  3. marn-0.1.0/README.md +101 -0
  4. marn-0.1.0/pyproject.toml +58 -0
  5. marn-0.1.0/src/marn/__init__.py +154 -0
  6. marn-0.1.0/src/marn/callbacks/__init__.py +7 -0
  7. marn-0.1.0/src/marn/callbacks/base.py +51 -0
  8. marn-0.1.0/src/marn/callbacks/early_stopping.py +60 -0
  9. marn-0.1.0/src/marn/callbacks/logger.py +47 -0
  10. marn-0.1.0/src/marn/checkpoint/__init__.py +12 -0
  11. marn-0.1.0/src/marn/checkpoint/load.py +174 -0
  12. marn-0.1.0/src/marn/checkpoint/save.py +122 -0
  13. marn-0.1.0/src/marn/checkpoint/schema.py +54 -0
  14. marn-0.1.0/src/marn/config/__init__.py +61 -0
  15. marn-0.1.0/src/marn/config/loss.py +50 -0
  16. marn-0.1.0/src/marn/config/model.py +67 -0
  17. marn-0.1.0/src/marn/config/trainer.py +39 -0
  18. marn-0.1.0/src/marn/distributed/__init__.py +18 -0
  19. marn-0.1.0/src/marn/distributed/ddp.py +180 -0
  20. marn-0.1.0/src/marn/distributed/memory.py +157 -0
  21. marn-0.1.0/src/marn/generators/__init__.py +22 -0
  22. marn-0.1.0/src/marn/generators/base.py +47 -0
  23. marn-0.1.0/src/marn/generators/finetuning.py +129 -0
  24. marn-0.1.0/src/marn/generators/grouped.py +138 -0
  25. marn-0.1.0/src/marn/generators/layer_generator.py +31 -0
  26. marn-0.1.0/src/marn/generators/layerwise.py +36 -0
  27. marn-0.1.0/src/marn/generators/lazy.py +105 -0
  28. marn-0.1.0/src/marn/generators/lrd.py +157 -0
  29. marn-0.1.0/src/marn/generators/single_vector.py +91 -0
  30. marn-0.1.0/src/marn/losses/__init__.py +23 -0
  31. marn-0.1.0/src/marn/losses/alignment.py +87 -0
  32. marn-0.1.0/src/marn/losses/base.py +23 -0
  33. marn-0.1.0/src/marn/losses/context.py +40 -0
  34. marn-0.1.0/src/marn/losses/mapping_loss.py +226 -0
  35. marn-0.1.0/src/marn/losses/outputs.py +24 -0
  36. marn-0.1.0/src/marn/losses/smoothness.py +147 -0
  37. marn-0.1.0/src/marn/losses/stability.py +79 -0
  38. marn-0.1.0/src/marn/losses/task.py +79 -0
  39. marn-0.1.0/src/marn/mappers/__init__.py +6 -0
  40. marn-0.1.0/src/marn/mappers/base.py +28 -0
  41. marn-0.1.0/src/marn/mappers/mlp_mapper.py +123 -0
  42. marn-0.1.0/src/marn/models/__init__.py +7 -0
  43. marn-0.1.0/src/marn/models/forward_result.py +29 -0
  44. marn-0.1.0/src/marn/models/mapping_model.py +142 -0
  45. marn-0.1.0/src/marn/models/target_model.py +80 -0
  46. marn-0.1.0/src/marn/modulation/__init__.py +8 -0
  47. marn-0.1.0/src/marn/modulation/additive.py +25 -0
  48. marn-0.1.0/src/marn/modulation/affine.py +29 -0
  49. marn-0.1.0/src/marn/modulation/base.py +16 -0
  50. marn-0.1.0/src/marn/modulation/low_rank.py +34 -0
  51. marn-0.1.0/src/marn/py.typed +1 -0
  52. marn-0.1.0/src/marn/registry.py +111 -0
  53. marn-0.1.0/src/marn/runtime/__init__.py +13 -0
  54. marn-0.1.0/src/marn/runtime/functional.py +41 -0
  55. marn-0.1.0/src/marn/runtime/lazy.py +58 -0
  56. marn-0.1.0/src/marn/runtime/parameter_spec.py +160 -0
  57. marn-0.1.0/src/marn/runtime/parameter_tree.py +54 -0
  58. marn-0.1.0/src/marn/strategies/__init__.py +17 -0
  59. marn-0.1.0/src/marn/strategies/base.py +14 -0
  60. marn-0.1.0/src/marn/strategies/finetuning.py +58 -0
  61. marn-0.1.0/src/marn/strategies/grouped.py +18 -0
  62. marn-0.1.0/src/marn/strategies/layerwise.py +15 -0
  63. marn-0.1.0/src/marn/strategies/lrd.py +41 -0
  64. marn-0.1.0/src/marn/strategies/slvt.py +24 -0
  65. marn-0.1.0/src/marn/trainers/__init__.py +17 -0
  66. marn-0.1.0/src/marn/trainers/batch_adapter.py +76 -0
  67. marn-0.1.0/src/marn/trainers/lr_finder.py +374 -0
  68. marn-0.1.0/src/marn/trainers/trainer.py +464 -0
marn-0.1.0/LICENSE ADDED
@@ -0,0 +1,161 @@
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
6
+
7
+ 1. Definitions.
8
+
9
+ "License" shall mean the terms and conditions for use, reproduction, and
10
+ distribution as defined by Sections 1 through 9 of this document.
11
+
12
+ "Licensor" shall mean the copyright owner or entity authorized by the
13
+ copyright owner that is granting the License.
14
+
15
+ "Legal Entity" shall mean the union of the acting entity and all other
16
+ entities that control, are controlled by, or are under common control with
17
+ that entity. For the purposes of this definition, "control" means (i) the
18
+ power, direct or indirect, to cause the direction or management of such
19
+ entity, whether by contract or otherwise, or (ii) ownership of fifty percent
20
+ (50%) or more of the outstanding shares, or (iii) beneficial ownership of
21
+ such entity.
22
+
23
+ "You" (or "Your") shall mean an individual or Legal Entity exercising
24
+ permissions granted by this License.
25
+
26
+ "Source" form shall mean the preferred form for making modifications,
27
+ including but not limited to software source code, documentation source, and
28
+ configuration files.
29
+
30
+ "Object" form shall mean any form resulting from mechanical transformation
31
+ or translation of a Source form, including but not limited to compiled object
32
+ code, generated documentation, and conversions to other media types.
33
+
34
+ "Work" shall mean the work of authorship, whether in Source or Object form,
35
+ made available under the License, as indicated by a copyright notice that is
36
+ included in or attached to the work.
37
+
38
+ "Derivative Works" shall mean any work, whether in Source or Object form,
39
+ that is based on (or derived from) the Work and for which the editorial
40
+ revisions, annotations, elaborations, or other modifications represent, as a
41
+ whole, an original work of authorship. For the purposes of this License,
42
+ Derivative Works shall not include works that remain separable from, or
43
+ merely link (or bind by name) to the interfaces of, the Work and Derivative
44
+ Works thereof.
45
+
46
+ "Contribution" shall mean any work of authorship, including the original
47
+ version of the Work and any modifications or additions to that Work or
48
+ Derivative Works thereof, that is intentionally submitted to Licensor for
49
+ inclusion in the Work by the copyright owner or by an individual or Legal
50
+ Entity authorized to submit on behalf of the copyright owner. For the purposes
51
+ of this definition, "submitted" means any form of electronic, verbal, or
52
+ written communication sent to the Licensor or its representatives, including
53
+ but not limited to communication on electronic mailing lists, source code
54
+ control systems, and issue tracking systems that are managed by, or on behalf
55
+ of, the Licensor for the purpose of discussing and improving the Work, but
56
+ excluding communication that is conspicuously marked or otherwise designated
57
+ in writing by the copyright owner as "Not a Contribution."
58
+
59
+ "Contributor" shall mean Licensor and any individual or Legal Entity on behalf
60
+ of whom a Contribution has been received by Licensor and subsequently
61
+ incorporated within the Work.
62
+
63
+ 2. Grant of Copyright License. Subject to the terms and conditions of this
64
+ License, each Contributor hereby grants to You a perpetual, worldwide,
65
+ non-exclusive, no-charge, royalty-free, irrevocable copyright license to
66
+ reproduce, prepare Derivative Works of, publicly display, publicly perform,
67
+ sublicense, and distribute the Work and such Derivative Works in Source or
68
+ Object form.
69
+
70
+ 3. Grant of Patent License. Subject to the terms and conditions of this
71
+ License, each Contributor hereby grants to You a perpetual, worldwide,
72
+ non-exclusive, no-charge, royalty-free, irrevocable (except as stated in this
73
+ section) patent license to make, have made, use, offer to sell, sell, import,
74
+ and otherwise transfer the Work, where such license applies only to those
75
+ patent claims licensable by such Contributor that are necessarily infringed by
76
+ their Contribution(s) alone or by combination of their Contribution(s) with
77
+ the Work to which such Contribution(s) was submitted. If You institute patent
78
+ litigation against any entity (including a cross-claim or counterclaim in a
79
+ lawsuit) alleging that the Work or a Contribution incorporated within the Work
80
+ constitutes direct or contributory patent infringement, then any patent
81
+ licenses granted to You under this License for that Work shall terminate as of
82
+ the date such litigation is filed.
83
+
84
+ 4. Redistribution. You may reproduce and distribute copies of the Work or
85
+ Derivative Works thereof in any medium, with or without modifications, and in
86
+ Source or Object form, provided that You meet the following conditions:
87
+
88
+ (a) You must give any other recipients of the Work or Derivative Works a copy
89
+ of this License; and
90
+
91
+ (b) You must cause any modified files to carry prominent notices stating that
92
+ You changed the files; and
93
+
94
+ (c) You must retain, in the Source form of any Derivative Works that You
95
+ distribute, all copyright, patent, trademark, and attribution notices from the
96
+ Source form of the Work, excluding those notices that do not pertain to any
97
+ part of the Derivative Works; and
98
+
99
+ (d) If the Work includes a "NOTICE" text file as part of its distribution,
100
+ then any Derivative Works that You distribute must include a readable copy of
101
+ the attribution notices contained within such NOTICE file, excluding those
102
+ notices that do not pertain to any part of the Derivative Works, in at least
103
+ one of the following places: within a NOTICE text file distributed as part of
104
+ the Derivative Works; within the Source form or documentation, if provided
105
+ along with the Derivative Works; or within a display generated by the
106
+ Derivative Works, if and wherever such third-party notices normally appear.
107
+ The contents of the NOTICE file are for informational purposes only and do not
108
+ modify the License. You may add Your own attribution notices within Derivative
109
+ Works that You distribute, alongside or as an addendum to the NOTICE text from
110
+ the Work, provided that such additional attribution notices cannot be construed
111
+ as modifying the License.
112
+
113
+ You may add Your own copyright statement to Your modifications and may provide
114
+ additional or different license terms and conditions for use, reproduction, or
115
+ distribution of Your modifications, or for any such Derivative Works as a
116
+ whole, provided Your use, reproduction, and distribution of the Work otherwise
117
+ complies with the conditions stated in this License.
118
+
119
+ 5. Submission of Contributions. Unless You explicitly state otherwise, any
120
+ Contribution intentionally submitted for inclusion in the Work by You to the
121
+ Licensor shall be under the terms and conditions of this License, without any
122
+ additional terms or conditions. Notwithstanding the above, nothing herein
123
+ shall supersede or modify the terms of any separate license agreement you may
124
+ have executed with Licensor regarding such Contributions.
125
+
126
+ 6. Trademarks. This License does not grant permission to use the trade names,
127
+ trademarks, service marks, or product names of the Licensor, except as
128
+ required for reasonable and customary use in describing the origin of the Work
129
+ and reproducing the content of the NOTICE file.
130
+
131
+ 7. Disclaimer of Warranty. Unless required by applicable law or agreed to in
132
+ writing, Licensor provides the Work (and each Contributor provides its
133
+ Contributions) on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
134
+ KIND, either express or implied, including, without limitation, any warranties
135
+ or conditions of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
136
+ PARTICULAR PURPOSE. You are solely responsible for determining the
137
+ appropriateness of using or redistributing the Work and assume any risks
138
+ associated with Your exercise of permissions under this License.
139
+
140
+ 8. Limitation of Liability. In no event and under no legal theory, whether in
141
+ tort (including negligence), contract, or otherwise, unless required by
142
+ applicable law (such as deliberate and grossly negligent acts) or agreed to in
143
+ writing, shall any Contributor be liable to You for damages, including any
144
+ direct, indirect, special, incidental, or consequential damages of any
145
+ character arising as a result of this License or out of the use or inability
146
+ to use the Work (including but not limited to damages for loss of goodwill,
147
+ work stoppage, computer failure or malfunction, or any and all other
148
+ commercial damages or losses), even if such Contributor has been advised of
149
+ the possibility of such damages.
150
+
151
+ 9. Accepting Warranty or Additional Liability. While redistributing the Work
152
+ or Derivative Works thereof, You may choose to offer, and charge a fee for,
153
+ acceptance of support, warranty, indemnity, or other liability obligations
154
+ and/or rights consistent with this License. However, in accepting such
155
+ obligations, You may act only on Your own behalf and on Your sole
156
+ responsibility, not on behalf of any other Contributor, and only if You agree
157
+ to indemnify, defend, and hold each Contributor harmless for any liability
158
+ incurred by, or claims asserted against, such Contributor by reason of your
159
+ accepting any such warranty or additional liability.
160
+
161
+ END OF TERMS AND CONDITIONS
marn-0.1.0/PKG-INFO ADDED
@@ -0,0 +1,125 @@
1
+ Metadata-Version: 2.4
2
+ Name: marn
3
+ Version: 0.1.0
4
+ Summary:
5
+ License: Apache-2.0
6
+ License-File: LICENSE
7
+ Author: Arjun Manjunath
8
+ Author-email: dev.arjunmnath@gmail.com
9
+ Requires-Python: >=3.12, !=2.7.*, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*, !=3.4.*, !=3.5.*, !=3.6.*, !=3.7.*, !=3.8.*, !=3.9.*, !=3.10.*, !=3.11.*
10
+ Classifier: License :: OSI Approved :: Apache Software License
11
+ Classifier: Programming Language :: Python :: 3
12
+ Classifier: Programming Language :: Python :: 3.12
13
+ Classifier: Programming Language :: Python :: 3.13
14
+ Classifier: Programming Language :: Python :: 3.14
15
+ Requires-Dist: accelerate (>=0.30)
16
+ Requires-Dist: matplotlib (>=3.11.2,<4.0.0)
17
+ Requires-Dist: pydantic (>=2.0)
18
+ Requires-Dist: pyyaml (>=6.0.3,<7.0.0)
19
+ Requires-Dist: scikit-learn (>=1.9.0,<2.0.0)
20
+ Requires-Dist: torch (>=2.6,<3.0)
21
+ Requires-Dist: torchvision (>=0.27.1,<0.28.0)
22
+ Description-Content-Type: text/markdown
23
+
24
+ # Manifold Regularized Networks (MaRN)
25
+
26
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue.svg)](https://opensource.org/licenses/Apache-2.0)
27
+ [![Python Version](https://img.shields.io/badge/python-3.12-blue)](https://www.python.org/downloads/)
28
+ [![PyTorch Version](https://img.shields.io/badge/PyTorch-%3E%3D2.6%20%7C%20%3C3.0-orange)](https://pytorch.org/)
29
+ [![Build Status](https://img.shields.io/badge/build-passing-brightgreen)](#)
30
+ [![Documentation Status](https://readthedocs.org/projects/mapping-networks/badge/?version=latest)](https://mapping-networks.readthedocs.io/en/latest/?badge=latest)
31
+ [![arXiv](https://img.shields.io/badge/arXiv-2602.19134-b31b1b.svg)](https://arxiv.org/abs/2602.19134)
32
+
33
+ `marn` is a model-agnostic, production-oriented PyTorch package for training target models through low-dimensional parameter manifolds. It is inspired by the paper [**Mapping Networks**](https://doi.org/10.48550/arXiv.2602.19134). The library decouples model architecture from parameter representation, allowing you to optimize neural networks by updating compact, trainable latent coordinates instead of mutating the target module's parameters.
34
+
35
+ ---
36
+
37
+ ## Installation
38
+
39
+ Install `marn` from your local workspace:
40
+
41
+ ```bash
42
+ pip install marn
43
+ ```
44
+
45
+ Or add it using Poetry:
46
+
47
+ ```bash
48
+ poetry add marn
49
+ ```
50
+
51
+ ---
52
+
53
+ ## Quickstart
54
+
55
+ Train a standard PyTorch model through a low-dimensional layer-wise latent representation:
56
+
57
+ ```python
58
+ import torch
59
+ from torch import nn
60
+ from torch.utils.data import DataLoader, TensorDataset
61
+ from marn import MappingModel, MappingLoss, ClassificationLoss, MappingTrainer
62
+
63
+ # 1. Create target model and training data
64
+ target = nn.Sequential(nn.Linear(10, 16), nn.ReLU(), nn.Linear(16, 2))
65
+ data = TensorDataset(torch.randn(100, 10), torch.randint(0, 2, (100,)))
66
+ loader = DataLoader(data, batch_size=16)
67
+
68
+ # 2. Wrap target model with low-dimensional mapping strategy
69
+ model = MappingModel(
70
+ target_model=target,
71
+ latent_dim=32,
72
+ strategy="layerwise"
73
+ )
74
+
75
+ # 3. Configure composite loss and trainer
76
+ loss_fn = MappingLoss(task_loss=ClassificationLoss())
77
+ trainer = MappingTrainer(
78
+ model=model,
79
+ train_loader=loader,
80
+ loss_fn=loss_fn,
81
+ learning_rate=1e-3,
82
+ )
83
+
84
+ # 4. Train the latent parameters
85
+ trainer.fit(epochs=5)
86
+ ```
87
+
88
+ ---
89
+
90
+ ## Core Concepts
91
+
92
+ - **Latent Manifold**: Trainable coordinates $z$ which undergo transformation to map to target parameter space.
93
+ - **BaseMapper**: Projects $z$ to flat generated parameter descriptors (e.g., using fixed, orthogonal projections).
94
+ - **BaseModulation**: Integrates generated descriptors back into target parameters (e.g., additive $W_{ij} \leftarrow W_{ij} + \alpha z_i$, or affine).
95
+ - **Generation Strategy**: Defines the mapping scope.
96
+ - `"slvt"` (Single Latent Vector Training) projects the entire model from a single global latent vector.
97
+ - `"layerwise"` constructs independent smaller latent coordinates per layer.
98
+ - `"grouped"` allows custom parameter subdivision.
99
+ - **MappingLoss**: A composite loss that balances task loss (classification/regression) with stability, smoothness, and cosine alignment regularization components.
100
+
101
+ ---
102
+
103
+ ## Paper Correspondence
104
+
105
+ The code maps directly to the concepts defined in the paper:
106
+ - **Fixed Projection Matrices**: Registered as PyTorch buffers inside `MLPMapper` so they stay frozen and are excluded from DDP/optimizer updates.
107
+ - **Additive Modulation**: Implemented in `AdditiveModulation` representing $w_{ij} \leftarrow w_{ij} + \alpha \cdot z_i$.
108
+ - **Regularization Terms**: Fully implemented in `MappingLoss`:
109
+ - **Stability Loss** ($L_{\text{stability}}$): penalizes changes in output predictions when adding small noise perturbation to the latent vector (`StabilityLoss`).
110
+ - **Smoothness Loss** ($L_{\text{smoothness}}$): penalizes the Jacobian norm of the mapper to enforce a smooth manifold (`SmoothnessLoss`).
111
+ - **Alignment Loss** ($L_{\text{alignment}}$): maximizes alignment via cosine distance between latent vectors and weight summaries (`AlignmentLoss`).
112
+
113
+ ---
114
+
115
+ ## Documentation & cookbook
116
+
117
+ - **Docs**: [mapping-networks.readthedocs.io](https://mapping-networks.readthedocs.io/en/latest/)
118
+ - **User guide & API**: organized under the PyData Sphinx theme in `docs/`
119
+ - **Cookbook**: eight runnable scripts in [`cookbook/`](cookbook/) (also summarized in the [Cookbook docs](https://mapping-networks.readthedocs.io/en/latest/cookbook/index.html))
120
+
121
+ ```bash
122
+ poetry run python cookbook/01_single_latent_classification.py
123
+ ```
124
+
125
+
marn-0.1.0/README.md ADDED
@@ -0,0 +1,101 @@
1
+ # Manifold Regularized Networks (MaRN)
2
+
3
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue.svg)](https://opensource.org/licenses/Apache-2.0)
4
+ [![Python Version](https://img.shields.io/badge/python-3.12-blue)](https://www.python.org/downloads/)
5
+ [![PyTorch Version](https://img.shields.io/badge/PyTorch-%3E%3D2.6%20%7C%20%3C3.0-orange)](https://pytorch.org/)
6
+ [![Build Status](https://img.shields.io/badge/build-passing-brightgreen)](#)
7
+ [![Documentation Status](https://readthedocs.org/projects/mapping-networks/badge/?version=latest)](https://mapping-networks.readthedocs.io/en/latest/?badge=latest)
8
+ [![arXiv](https://img.shields.io/badge/arXiv-2602.19134-b31b1b.svg)](https://arxiv.org/abs/2602.19134)
9
+
10
+ `marn` is a model-agnostic, production-oriented PyTorch package for training target models through low-dimensional parameter manifolds. It is inspired by the paper [**Mapping Networks**](https://doi.org/10.48550/arXiv.2602.19134). The library decouples model architecture from parameter representation, allowing you to optimize neural networks by updating compact, trainable latent coordinates instead of mutating the target module's parameters.
11
+
12
+ ---
13
+
14
+ ## Installation
15
+
16
+ Install `marn` from your local workspace:
17
+
18
+ ```bash
19
+ pip install marn
20
+ ```
21
+
22
+ Or add it using Poetry:
23
+
24
+ ```bash
25
+ poetry add marn
26
+ ```
27
+
28
+ ---
29
+
30
+ ## Quickstart
31
+
32
+ Train a standard PyTorch model through a low-dimensional layer-wise latent representation:
33
+
34
+ ```python
35
+ import torch
36
+ from torch import nn
37
+ from torch.utils.data import DataLoader, TensorDataset
38
+ from marn import MappingModel, MappingLoss, ClassificationLoss, MappingTrainer
39
+
40
+ # 1. Create target model and training data
41
+ target = nn.Sequential(nn.Linear(10, 16), nn.ReLU(), nn.Linear(16, 2))
42
+ data = TensorDataset(torch.randn(100, 10), torch.randint(0, 2, (100,)))
43
+ loader = DataLoader(data, batch_size=16)
44
+
45
+ # 2. Wrap target model with low-dimensional mapping strategy
46
+ model = MappingModel(
47
+ target_model=target,
48
+ latent_dim=32,
49
+ strategy="layerwise"
50
+ )
51
+
52
+ # 3. Configure composite loss and trainer
53
+ loss_fn = MappingLoss(task_loss=ClassificationLoss())
54
+ trainer = MappingTrainer(
55
+ model=model,
56
+ train_loader=loader,
57
+ loss_fn=loss_fn,
58
+ learning_rate=1e-3,
59
+ )
60
+
61
+ # 4. Train the latent parameters
62
+ trainer.fit(epochs=5)
63
+ ```
64
+
65
+ ---
66
+
67
+ ## Core Concepts
68
+
69
+ - **Latent Manifold**: Trainable coordinates $z$ which undergo transformation to map to target parameter space.
70
+ - **BaseMapper**: Projects $z$ to flat generated parameter descriptors (e.g., using fixed, orthogonal projections).
71
+ - **BaseModulation**: Integrates generated descriptors back into target parameters (e.g., additive $W_{ij} \leftarrow W_{ij} + \alpha z_i$, or affine).
72
+ - **Generation Strategy**: Defines the mapping scope.
73
+ - `"slvt"` (Single Latent Vector Training) projects the entire model from a single global latent vector.
74
+ - `"layerwise"` constructs independent smaller latent coordinates per layer.
75
+ - `"grouped"` allows custom parameter subdivision.
76
+ - **MappingLoss**: A composite loss that balances task loss (classification/regression) with stability, smoothness, and cosine alignment regularization components.
77
+
78
+ ---
79
+
80
+ ## Paper Correspondence
81
+
82
+ The code maps directly to the concepts defined in the paper:
83
+ - **Fixed Projection Matrices**: Registered as PyTorch buffers inside `MLPMapper` so they stay frozen and are excluded from DDP/optimizer updates.
84
+ - **Additive Modulation**: Implemented in `AdditiveModulation` representing $w_{ij} \leftarrow w_{ij} + \alpha \cdot z_i$.
85
+ - **Regularization Terms**: Fully implemented in `MappingLoss`:
86
+ - **Stability Loss** ($L_{\text{stability}}$): penalizes changes in output predictions when adding small noise perturbation to the latent vector (`StabilityLoss`).
87
+ - **Smoothness Loss** ($L_{\text{smoothness}}$): penalizes the Jacobian norm of the mapper to enforce a smooth manifold (`SmoothnessLoss`).
88
+ - **Alignment Loss** ($L_{\text{alignment}}$): maximizes alignment via cosine distance between latent vectors and weight summaries (`AlignmentLoss`).
89
+
90
+ ---
91
+
92
+ ## Documentation & cookbook
93
+
94
+ - **Docs**: [mapping-networks.readthedocs.io](https://mapping-networks.readthedocs.io/en/latest/)
95
+ - **User guide & API**: organized under the PyData Sphinx theme in `docs/`
96
+ - **Cookbook**: eight runnable scripts in [`cookbook/`](cookbook/) (also summarized in the [Cookbook docs](https://mapping-networks.readthedocs.io/en/latest/cookbook/index.html))
97
+
98
+ ```bash
99
+ poetry run python cookbook/01_single_latent_classification.py
100
+ ```
101
+
@@ -0,0 +1,58 @@
1
+ [tool.poetry]
2
+ name = "marn"
3
+ version = "0.1.0"
4
+ description = ""
5
+ authors = ["Arjun Manjunath <dev.arjunmnath@gmail.com>"]
6
+ readme = "README.md"
7
+ license = "Apache-2.0"
8
+
9
+ packages = [
10
+ { include = "marn", from = "src" }
11
+ ]
12
+
13
+ [tool.poetry.dependencies]
14
+ python = ">=3.12,<3.14.1 || >3.14.1,<3.15"
15
+ torch = ">=2.6,<3.0"
16
+ pydantic = ">=2.0"
17
+ pyyaml = "^6.0.3"
18
+ torchvision = "^0.27.1"
19
+ accelerate = ">=0.30"
20
+ scikit-learn = "^1.9.0"
21
+ matplotlib = "^3.11.2"
22
+
23
+ [tool.poetry.group.dev.dependencies]
24
+ pytest = "^8.3"
25
+ pytest-cov = "^6.0"
26
+ ruff = "^0.12"
27
+ mypy = "^1.16"
28
+ pre-commit = "^4.2"
29
+ types-pyyaml = "^6.0.12.20260518"
30
+
31
+
32
+ [tool.poetry.group.docs.dependencies]
33
+ sphinx = "^9.1.0"
34
+ myst-parser = "^5.1.0"
35
+ pydata-sphinx-theme = "^0.15.2"
36
+ sphinx-autobuild = "^2025.8.25"
37
+
38
+
39
+ [tool.ruff]
40
+ line-length = 100
41
+
42
+ [tool.pytest.ini_options]
43
+ testpaths = ["tests"]
44
+
45
+ [tool.mypy]
46
+ strict = true
47
+
48
+ [[tool.mypy.overrides]]
49
+ module = [
50
+ "matplotlib.*",
51
+ "torchvision.*",
52
+ "accelerate.*",
53
+ ]
54
+ ignore_missing_imports = true
55
+
56
+ [build-system]
57
+ requires = ["poetry-core"]
58
+ build-backend = "poetry.core.masonry.api"
@@ -0,0 +1,154 @@
1
+ """Train PyTorch models through low-dimensional parameter mappings."""
2
+
3
+ from marn.generators import (
4
+ GroupedGenerator,
5
+ LayerGenerator,
6
+ LazyLayerwiseGenerator,
7
+ LayerwiseGenerator,
8
+ ParameterGenerator,
9
+ SingleVectorGenerator,
10
+ FineTuningGenerator,
11
+ LRDGenerator,
12
+ )
13
+ from marn.distributed import (
14
+ benchmark_strategies,
15
+ cleanup_ddp,
16
+ is_ddp_available,
17
+ profile_peak_memory,
18
+ setup_ddp,
19
+ wrap_ddp,
20
+ )
21
+ from marn.models import (
22
+ ForwardResult,
23
+ MappingModel,
24
+ TargetModel,
25
+ UnsupportedTargetModelError,
26
+ )
27
+ from marn.losses import (
28
+ AlignmentLoss,
29
+ BaseLoss,
30
+ ClassificationLoss,
31
+ LossOutput,
32
+ MappingLoss,
33
+ RegressionLoss,
34
+ SmoothnessLoss,
35
+ StabilityLoss,
36
+ TaskLoss,
37
+ TrainingContext,
38
+ )
39
+ from marn.mappers import BaseMapper, MLPMapper, ResidualMLPMapper
40
+ from marn.modulation import (
41
+ AdditiveModulation,
42
+ AffineModulation,
43
+ BaseModulation,
44
+ LowRankModulation,
45
+ )
46
+ from marn.runtime import ParameterEntry, ParameterSpec, ParameterTree
47
+ from marn.strategies import (
48
+ GroupedStrategy,
49
+ LayerwiseStrategy,
50
+ SLVTStrategy,
51
+ FineTuningStrategy,
52
+ LRDStrategy,
53
+ )
54
+ from marn.config import (
55
+ GeneratorConfig,
56
+ LossConfig,
57
+ MapperConfig,
58
+ MappingConfig,
59
+ TaskLossConfig,
60
+ TrainerConfig,
61
+ load_config,
62
+ )
63
+ from marn.registry import (
64
+ GENERATOR_REGISTRY,
65
+ LOSS_REGISTRY,
66
+ MAPPER_REGISTRY,
67
+ MODULATION_REGISTRY,
68
+ Registry,
69
+ )
70
+ from marn.callbacks import Callback, EarlyStopping, MetricLogger
71
+ from marn.trainers import (
72
+ BatchAdapter,
73
+ MappingBatchAdapter,
74
+ MappingTrainer,
75
+ TupleBatchAdapter,
76
+ LRFinderResult,
77
+ )
78
+ from marn.checkpoint import (
79
+ CheckpointCompatibilityError,
80
+ CheckpointSchema,
81
+ load_checkpoint,
82
+ save_checkpoint,
83
+ )
84
+
85
+ __all__ = [
86
+ "AdditiveModulation",
87
+ "AffineModulation",
88
+ "AlignmentLoss",
89
+ "BaseLoss",
90
+ "BaseMapper",
91
+ "BaseModulation",
92
+ "BatchAdapter",
93
+ "Callback",
94
+ "CheckpointCompatibilityError",
95
+ "CheckpointSchema",
96
+ "ClassificationLoss",
97
+ "EarlyStopping",
98
+ "ForwardResult",
99
+ "FineTuningGenerator",
100
+ "GENERATOR_REGISTRY",
101
+ "GeneratorConfig",
102
+ "GroupedGenerator",
103
+ "GroupedStrategy",
104
+ "LOSS_REGISTRY",
105
+ "LayerGenerator",
106
+ "LazyLayerwiseGenerator",
107
+ "LayerwiseGenerator",
108
+ "LayerwiseStrategy",
109
+ "LRDGenerator",
110
+ "LossConfig",
111
+ "LossOutput",
112
+ "LowRankModulation",
113
+ "MAPPER_REGISTRY",
114
+ "MODULATION_REGISTRY",
115
+ "MapperConfig",
116
+ "MappingBatchAdapter",
117
+ "MappingConfig",
118
+ "MappingLoss",
119
+ "MappingModel",
120
+ "MappingTrainer",
121
+ "MLPMapper",
122
+ "MetricLogger",
123
+ "ParameterEntry",
124
+ "ParameterGenerator",
125
+ "ParameterSpec",
126
+ "ParameterTree",
127
+ "RegressionLoss",
128
+ "Registry",
129
+ "ResidualMLPMapper",
130
+ "SLVTStrategy",
131
+ "SingleVectorGenerator",
132
+ "FineTuningStrategy",
133
+ "LRDStrategy",
134
+ "LRFinderResult",
135
+ "SmoothnessLoss",
136
+ "StabilityLoss",
137
+ "TargetModel",
138
+ "TaskLoss",
139
+ "TaskLossConfig",
140
+ "TrainerConfig",
141
+ "TrainingContext",
142
+ "TupleBatchAdapter",
143
+ "UnsupportedTargetModelError",
144
+ "benchmark_strategies",
145
+ "cleanup_ddp",
146
+ "is_ddp_available",
147
+ "load_checkpoint",
148
+ "load_config",
149
+ "profile_peak_memory",
150
+ "save_checkpoint",
151
+ "setup_ddp",
152
+ "wrap_ddp",
153
+ ]
154
+ __version__ = "0.1.0"
@@ -0,0 +1,7 @@
1
+ """Training lifecycle callbacks."""
2
+
3
+ from marn.callbacks.base import Callback
4
+ from marn.callbacks.early_stopping import EarlyStopping
5
+ from marn.callbacks.logger import MetricLogger
6
+
7
+ __all__ = ["Callback", "EarlyStopping", "MetricLogger"]
@@ -0,0 +1,51 @@
1
+ """Observer callback base class for training lifecycle hooks."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import TYPE_CHECKING, Any
6
+
7
+ if TYPE_CHECKING:
8
+ from marn.losses.outputs import LossOutput
9
+
10
+
11
+ class Callback:
12
+ """Base class for training lifecycle callbacks.
13
+
14
+ All hooks are no-ops by default — subclass and override the ones you need.
15
+ Callbacks receive the trainer instance so they can inspect or modify state
16
+ (e.g. set ``trainer.should_stop = True`` for early termination).
17
+
18
+ Hook execution order within each event:
19
+
20
+ 1. ``on_fit_start``
21
+ 2. For each epoch:
22
+ a. ``on_epoch_start``
23
+ b. For each batch: ``on_batch_start`` → ``on_batch_end``
24
+ c. ``on_validation_start`` → ``on_validation_end``
25
+ d. ``on_epoch_end``
26
+ 3. ``on_fit_end``
27
+ """
28
+
29
+ def on_fit_start(self, trainer: Any) -> None:
30
+ """Called once at the beginning of :meth:`fit`."""
31
+
32
+ def on_fit_end(self, trainer: Any) -> None:
33
+ """Called once at the end of :meth:`fit`."""
34
+
35
+ def on_epoch_start(self, trainer: Any, epoch: int) -> None:
36
+ """Called at the beginning of each epoch."""
37
+
38
+ def on_epoch_end(self, trainer: Any, epoch: int, metrics: dict[str, float]) -> None:
39
+ """Called at the end of each epoch with aggregated metrics."""
40
+
41
+ def on_batch_start(self, trainer: Any, batch_idx: int) -> None:
42
+ """Called before processing each batch."""
43
+
44
+ def on_batch_end(self, trainer: Any, batch_idx: int, loss_output: LossOutput) -> None:
45
+ """Called after processing each batch with the loss output."""
46
+
47
+ def on_validation_start(self, trainer: Any) -> None:
48
+ """Called before the validation loop."""
49
+
50
+ def on_validation_end(self, trainer: Any, metrics: dict[str, float]) -> None:
51
+ """Called after the validation loop with validation metrics."""