e3nn_mlx.models.gate_points_2102.Network¶
- class e3nn_mlx.models.gate_points_2102.Network(irreps_in, irreps_hidden, irreps_out, irreps_node_attr, irreps_edge_attr, layers, max_radius, number_of_basis, radial_layers, radial_neurons, num_neighbors, num_nodes, reduce_output=True)[source]¶
Gated E(3)-equivariant network operating on one or more point graphs.
- Parameters:
layers (int)
max_radius (float)
number_of_basis (int)
radial_layers (int)
radial_neurons (int)
num_neighbors (float)
num_nodes (float)
reduce_output (bool)
- __init__(irreps_in, irreps_hidden, irreps_out, irreps_node_attr, irreps_edge_attr, layers, max_radius, number_of_basis, radial_layers, radial_neurons, num_neighbors, num_nodes, reduce_output=True)[source]¶
Should be called by the subclasses of
Module.- Parameters:
layers (int)
max_radius (float)
number_of_basis (int)
radial_layers (int)
radial_neurons (int)
num_neighbors (float)
num_nodes (float)
reduce_output (bool)
- Return type:
None
Methods
__init__(irreps_in, irreps_hidden, ...[, ...])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_with_edges(positions, node_input, ...)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.