e3nn_mlx.o3.ReducedTensorProducts¶
- class e3nn_mlx.o3.ReducedTensorProducts(formula, filter_ir_out=None, filter_ir_mid=None, **irreps)[source]¶
- Parameters:
- __init__(formula, filter_ir_out=None, filter_ir_mid=None, **irreps)[source]¶
Should be called by the subclasses of
Module.
Methods
__init__(formula[, filter_ir_out, filter_ir_mid])Should be called by the subclasses of
Module.apply(map_fn[, filter_fn])Map all the parameters using the provided
map_fnand immediately update the module with the mapped parameters.apply_to_modules(apply_fn)Apply a function to all the modules in this instance (including this instance).
children()Return the direct descendants of this Module instance.
clear()copy()eval()Set the model to evaluation mode.
filter_and_map(filter_fn[, map_fn, is_leaf_fn])Recursively filter the contents of the module using
filter_fn, namely only select keys and values wherefilter_fnreturns true.forward(*inputs)Call self as a function.
freeze(*[, recurse, keys, strict])Freeze the Module's parameters or some of them.
fromkeys(iterable[, value])Create a new dictionary with keys from iterable and values set to value.
get(key[, default])Return the value for key if key is in the dictionary, else default.
is_module(value)items()keys()leaf_modules()Return the submodules that do not contain other modules.
load_weights(file_or_weights[, strict])Update the model's weights from a
.npz, a.safetensorsfile, or a list.modules()Return a list with all the modules in this instance.
named_modules()Return a list with all the modules in this instance and their name with dot notation.
parameters()Recursively return all the
mlx.core.arraymembers of this Module as a dict of dicts and lists.pop(k[,d])If the key is not found, return the default if given; otherwise, raise a KeyError.
popitem()Remove and return a (key, value) pair as a 2-tuple.
save_weights(file)Save the model's weights to a file.
set_dtype(dtype[, predicate])Set the dtype of the module's parameters.
setdefault(key[, default])Insert key with a value of default if key is not in the dictionary.
train([mode])Set the model in or out of training mode.
trainable_parameter_filter(module, key, value)trainable_parameters()Recursively return all the non frozen
mlx.core.arraymembers of this Module as a dict of dicts and lists.unfreeze(*[, recurse, keys, strict])Unfreeze the Module's parameters or some of them.
update(parameters[, strict])Replace the parameters of this Module with the provided ones in the dict of dicts and lists.
update_modules(modules[, strict])Replace the child modules of this
Moduleinstance with the provided ones in the dict of dicts and lists.valid_child_filter(module, key, value)valid_parameter_filter(module, key, value)values()Attributes
change_of_basisstateThe module's state dictionary
trainingBoolean indicating if the model is in training mode.