|
23 | 23 | from optax.transforms._accumulation import skip_not_finite |
24 | 24 | from optax.transforms._accumulation import trace |
25 | 25 | 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 |
26 | 41 | from optax.transforms._conditionality import apply_if_finite |
27 | 42 | from optax.transforms._conditionality import ApplyIfFiniteState |
28 | 43 | from optax.transforms._conditionality import conditionally_mask |
29 | 44 | from optax.transforms._conditionality import conditionally_transform |
30 | 45 | from optax.transforms._conditionality import ConditionallyMaskState |
31 | 46 | from optax.transforms._conditionality import ConditionallyTransformState |
32 | 47 | 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 |
33 | 52 | from optax.transforms._layouts import flatten |
34 | 53 | from optax.transforms._masking import masked |
35 | 54 | from optax.transforms._masking import MaskedNode |
36 | 55 | 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