| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| """kfac-jax public APIs.""" |
|
|
| from kfac_jax._src import curvature_blocks |
| from kfac_jax._src import curvature_estimator |
| from kfac_jax._src import layers_and_loss_tags |
| from kfac_jax._src import loss_functions |
| from kfac_jax._src import optimizer |
| from kfac_jax._src import patches_second_moment |
| from kfac_jax._src import tag_graph_matcher |
| from kfac_jax._src import tracer |
| from kfac_jax._src import utils |
|
|
|
|
| __version__ = "0.0.8" |
|
|
| |
| patches_moments = patches_second_moment.patches_moments |
| patches_moments_explicit = patches_second_moment.patches_moments_explicit |
|
|
| |
| LayerData = layers_and_loss_tags.LayerData |
| LayerMetaData = layers_and_loss_tags.LayerMetaData |
| LossTag = layers_and_loss_tags.LossTag |
| LayerTag = layers_and_loss_tags.LayerTag |
| register_generic = layers_and_loss_tags.register_generic |
| register_dense = layers_and_loss_tags.register_dense |
| register_conv2d = layers_and_loss_tags.register_conv2d |
| register_scale_and_shift = layers_and_loss_tags.register_scale_and_shift |
|
|
| |
| auto_register_tags = tag_graph_matcher.auto_register_tags |
|
|
| |
| ProcessedJaxpr = tracer.ProcessedJaxpr |
| LayerVjpData = tracer.LayerVjpData |
| loss_tags_vjp = tracer.loss_tags_vjp |
| loss_tags_jvp = tracer.loss_tags_jvp |
| loss_tags_hvp = tracer.loss_tags_hvp |
| layer_tags_vjp = tracer.layer_tags_vjp |
|
|
| |
| LossFunction = loss_functions.LossFunction |
| NegativeLogProbLoss = loss_functions.NegativeLogProbLoss |
| DistributionNegativeLogProbLoss = loss_functions.DistributionNegativeLogProbLoss |
| NormalMeanNegativeLogProbLoss = loss_functions.NormalMeanNegativeLogProbLoss |
| NormalMeanVarianceNegativeLogProbLoss = ( |
| loss_functions.NormalMeanVarianceNegativeLogProbLoss) |
| MultiBernoulliNegativeLogProbLoss = ( |
| loss_functions.MultiBernoulliNegativeLogProbLoss) |
| CategoricalLogitsNegativeLogProbLoss = ( |
| loss_functions.CategoricalLogitsNegativeLogProbLoss) |
| OneHotCategoricalLogitsNegativeLogProbLoss = ( |
| loss_functions.OneHotCategoricalLogitsNegativeLogProbLoss) |
| register_sigmoid_cross_entropy_loss = ( |
| loss_functions.register_sigmoid_cross_entropy_loss) |
| register_multi_bernoulli_predictive_distribution = ( |
| loss_functions.register_multi_bernoulli_predictive_distribution) |
| register_softmax_cross_entropy_loss = ( |
| loss_functions.register_softmax_cross_entropy_loss) |
| register_categorical_predictive_distribution = ( |
| loss_functions.register_categorical_predictive_distribution) |
| register_squared_error_loss = loss_functions.register_squared_error_loss |
| register_normal_predictive_distribution = ( |
| loss_functions.register_normal_predictive_distribution) |
|
|
| |
| CurvatureBlock = curvature_blocks.CurvatureBlock |
| ScaledIdentity = curvature_blocks.ScaledIdentity |
| Diagonal = curvature_blocks.Diagonal |
| Full = curvature_blocks.Full |
| KroneckerFactored = curvature_blocks.KroneckerFactored |
| NaiveDiagonal = curvature_blocks.NaiveDiagonal |
| NaiveFull = curvature_blocks.NaiveFull |
| NaiveTNT = curvature_blocks.NaiveTNT |
| DenseDiagonal = curvature_blocks.DenseDiagonal |
| DenseFull = curvature_blocks.DenseFull |
| DenseTwoKroneckerFactored = curvature_blocks.DenseTwoKroneckerFactored |
| RepeatedDenseKroneckerFactored = curvature_blocks.RepeatedDenseKroneckerFactored |
| DenseTNT = curvature_blocks.DenseTNT |
| Conv2DDiagonal = curvature_blocks.Conv2DDiagonal |
| Conv2DFull = curvature_blocks.Conv2DFull |
| Conv2DTwoKroneckerFactored = curvature_blocks.Conv2DTwoKroneckerFactored |
| Conv2DTNT = curvature_blocks.Conv2DTNT |
| ScaleAndShiftDiagonal = curvature_blocks.ScaleAndShiftDiagonal |
| ScaleAndShiftFull = curvature_blocks.ScaleAndShiftFull |
| set_max_parallel_elements = curvature_blocks.set_max_parallel_elements |
| get_max_parallel_elements = curvature_blocks.get_max_parallel_elements |
| set_default_eigen_decomposition_threshold = ( |
| curvature_blocks.set_default_eigen_decomposition_threshold) |
| get_default_eigen_decomposition_threshold = ( |
| curvature_blocks.get_default_eigen_decomposition_threshold) |
|
|
| |
| CurvatureEstimator = curvature_estimator.CurvatureEstimator |
| BlockDiagonalCurvature = curvature_estimator.BlockDiagonalCurvature |
| ExplicitExactCurvature = curvature_estimator.ExplicitExactCurvature |
| ImplicitExactCurvature = curvature_estimator.ImplicitExactCurvature |
| set_default_tag_to_block_ctor = ( |
| curvature_estimator.set_default_tag_to_block_ctor) |
| get_default_tag_to_block_ctor = ( |
| curvature_estimator.get_default_tag_to_block_ctor) |
| OptaxPreconditioner = curvature_estimator.OptaxPreconditioner |
| OptaxPreconditionState = curvature_estimator.OptaxPreconditionState |
|
|
| |
| Optimizer = optimizer.Optimizer |
|
|
| HAIKU_BIASES = optimizer.HAIKU_BIASES |
| HAIKU_BIASES_AND_NORMS = optimizer.HAIKU_BIASES_AND_NORMS |
|
|
| __all__ = ( |
| |
| "utils", |
| "patches_second_moment", |
| "layers_and_loss_tags", |
| "loss_functions", |
| "tag_graph_matcher", |
| "tracer", |
| "curvature_blocks", |
| "curvature_estimator", |
| "optimizer", |
| |
| "patches_moments", |
| "patches_moments_explicit", |
| |
| "LossTag", |
| "LayerTag", |
| "register_generic", |
| "register_dense", |
| "register_conv2d", |
| "register_scale_and_shift", |
| |
| "auto_register_tags", |
| |
| "ProcessedJaxpr", |
| "loss_tags_vjp", |
| "loss_tags_jvp", |
| "loss_tags_hvp", |
| "layer_tags_vjp", |
| |
| "LossFunction", |
| "NegativeLogProbLoss", |
| "DistributionNegativeLogProbLoss", |
| "NormalMeanNegativeLogProbLoss", |
| "NormalMeanVarianceNegativeLogProbLoss", |
| "MultiBernoulliNegativeLogProbLoss", |
| "CategoricalLogitsNegativeLogProbLoss", |
| "OneHotCategoricalLogitsNegativeLogProbLoss", |
| "register_sigmoid_cross_entropy_loss", |
| "register_multi_bernoulli_predictive_distribution", |
| "register_softmax_cross_entropy_loss", |
| "register_categorical_predictive_distribution", |
| "register_squared_error_loss", |
| "register_normal_predictive_distribution", |
| |
| "CurvatureBlock", |
| "ScaledIdentity", |
| "Diagonal", |
| "Full", |
| "KroneckerFactored", |
| "NaiveDiagonal", |
| "NaiveFull", |
| "NaiveTNT", |
| "DenseDiagonal", |
| "DenseFull", |
| "DenseTwoKroneckerFactored", |
| "RepeatedDenseKroneckerFactored", |
| "DenseTNT", |
| "Conv2DDiagonal", |
| "Conv2DFull", |
| "Conv2DTwoKroneckerFactored", |
| "Conv2DTNT", |
| "ScaleAndShiftDiagonal", |
| "ScaleAndShiftFull", |
| "set_max_parallel_elements", |
| "get_max_parallel_elements", |
| "set_default_eigen_decomposition_threshold", |
| "get_default_eigen_decomposition_threshold", |
| |
| "CurvatureEstimator", |
| "BlockDiagonalCurvature", |
| "ExplicitExactCurvature", |
| "ImplicitExactCurvature", |
| "set_default_tag_to_block_ctor", |
| "get_default_tag_to_block_ctor", |
| |
| "Optimizer", |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| try: |
| del _src |
| except NameError: |
| pass |
|
|