Skip to content

Commit 99587dc

Browse files
mtthssOptaxDev
authored andcommitted
Expose components in sub-package
PiperOrigin-RevId: 638188465
1 parent 36ee9f4 commit 99587dc

3 files changed

Lines changed: 60 additions & 0 deletions

File tree

optax/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
from optax import projections
2424
from optax import schedules
2525
from optax import second_order
26+
from optax import transforms
2627
from optax import tree_utils
2728
from optax._src.alias import adabelief
2829
from optax._src.alias import adadelta

optax/optax_test.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,14 +15,17 @@
1515
"""Tests for optax."""
1616

1717
from absl.testing import absltest
18+
1819
import optax
20+
from optax import transforms
1921

2022

2123
class OptaxTest(absltest.TestCase):
2224
"""Test optax can be imported correctly."""
2325

2426
def test_import(self):
2527
self.assertTrue(hasattr(optax, 'GradientTransformation'))
28+
self.assertTrue(hasattr(transforms, 'partition'))
2629

2730

2831
if __name__ == '__main__':

optax/transforms/__init__.py

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,14 +23,70 @@
2323
from optax.transforms._accumulation import skip_not_finite
2424
from optax.transforms._accumulation import trace
2525
from optax.transforms._accumulation import TraceState
26+
from optax.transforms._adding import add_decayed_weights
27+
from optax.transforms._adding import add_noise
28+
from optax.transforms._adding import AddNoiseState
29+
from optax.transforms._clipping import adaptive_grad_clip
30+
from optax.transforms._clipping import clip
31+
from optax.transforms._clipping import clip_by_block_rms
32+
from optax.transforms._clipping import clip_by_global_norm
33+
from optax.transforms._clipping import per_example_global_norm_clip
34+
from optax.transforms._clipping import per_example_layer_norm_clip
35+
from optax.transforms._clipping import unitwise_clip
36+
from optax.transforms._clipping import unitwise_norm
37+
from optax.transforms._combining import chain
38+
from optax.transforms._combining import named_chain
39+
from optax.transforms._combining import partition
40+
from optax.transforms._combining import PartitionState
2641
from optax.transforms._conditionality import apply_if_finite
2742
from optax.transforms._conditionality import ApplyIfFiniteState
2843
from optax.transforms._conditionality import conditionally_mask
2944
from optax.transforms._conditionality import conditionally_transform
3045
from optax.transforms._conditionality import ConditionallyMaskState
3146
from optax.transforms._conditionality import ConditionallyTransformState
3247
from optax.transforms._conditionality import ConditionFn
48+
from optax.transforms._constraining import keep_params_nonnegative
49+
from optax.transforms._constraining import NonNegativeParamsState
50+
from optax.transforms._constraining import zero_nans
51+
from optax.transforms._constraining import ZeroNansState
3352
from optax.transforms._layouts import flatten
3453
from optax.transforms._masking import masked
3554
from optax.transforms._masking import MaskedNode
3655
from optax.transforms._masking import MaskedState
56+
57+
58+
__all__ = (
59+
"adaptive_grad_clip",
60+
"add_decayed_weights",
61+
"add_noise",
62+
"AddNoiseState",
63+
"apply_if_finite",
64+
"ApplyIfFiniteState",
65+
"chain",
66+
"clip_by_block_rms",
67+
"clip_by_global_norm",
68+
"clip",
69+
"conditionally_mask",
70+
"ConditionallyMaskState",
71+
"conditionally_transform",
72+
"ConditionallyTransformState",
73+
"ema",
74+
"EmaState",
75+
"flatten",
76+
"keep_params_nonnegative",
77+
"masked",
78+
"MaskedState",
79+
"MultiSteps",
80+
"MultiStepsState",
81+
"named_chain",
82+
"NonNegativeParamsState",
83+
"partition",
84+
"PartitionState",
85+
"ShouldSkipUpdateFunction",
86+
"skip_large_updates",
87+
"skip_not_finite",
88+
"trace",
89+
"TraceState",
90+
"zero_nans",
91+
"ZeroNansState",
92+
)

0 commit comments

Comments
 (0)