e3nn_mlx.models.v2106.MessagePassing¶
- class e3nn_mlx.models.v2106.MessagePassing(irreps_node_sequence, irreps_node_attr, irreps_edge_attr, fc_neurons, num_neighbors, *, use_custom_kernel=True)[source]¶
Sequence of v2106 convolutions and parity-aware gates.
- Parameters:
num_neighbors (float)
use_custom_kernel (bool)
- __init__(irreps_node_sequence, irreps_node_attr, irreps_edge_attr, fc_neurons, num_neighbors, *, use_custom_kernel=True)[source]¶
Should be called by the subclasses of
Module.- Parameters:
num_neighbors (float)
use_custom_kernel (bool)
- Return type:
None
Methods
__init__(irreps_node_sequence, ...[, ...])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_arrays(node_features, node_attr, ...)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
stateThe module's state dictionary
trainingBoolean indicating if the model is in training mode.