cayleypy.models.ModelConfig
- class cayleypy.models.ModelConfig(model_type: str, input_size: int, num_classes_for_one_hot: int, layers_sizes: list[int], weights_kaggle_id: str | None = None, weights_path: str | None = None, n_outputs: int = 1, tokenizer_groups: list[list[int]] | None = None, graph_hash: str | None = None)[source]
Configuration used to describe ML model.
Fields n_outputs, tokenizer_groups and graph_hash describe capabilities added after the first version of this class. Their defaults describe a single-output model without tokenization, which is not tied to a particular graph, so configs written before these fields existed keep working.
- Parameters:
model_type – Type of the model, one of “MLP” (see
MlpModel) or “RESMLP” (seeResMlpModel).input_size – Number of elements in one state.
num_classes_for_one_hot – Number of distinct values one element of a state can take.
layers_sizes – Sizes of hidden layers (for “RESMLP” these are sizes of residual blocks).
weights_kaggle_id – Id of the Kaggle model with weights (optional).
weights_path – Path to the file with weights (optional).
n_outputs – Number of outputs of the model. 1 means the model predicts distance for the state it is applied to. n_generators means the model predicts distance for every child of that state (Q-model).
tokenizer_groups – Specification of how elements of a state are grouped into tokens, as a list of
[group_size, num_groups]pairs. For example,[[3, 20], [2, 30]]means 20 tokens of 3 elements followed by 30 tokens of 2 elements. None means the state is not tokenized.graph_hash – Hash of the graph this model was trained for, see
cayleypy.models.graph_hash().
- __init__(model_type: str, input_size: int, num_classes_for_one_hot: int, layers_sizes: list[int], weights_kaggle_id: str | None = None, weights_path: str | None = None, n_outputs: int = 1, tokenizer_groups: list[list[int]] | None = None, graph_hash: str | None = None) None
Methods
__init__(model_type, input_size, ...[, ...])Creates model described by this config, with randomly initialized weights.
from_dict(cfg)Creates config from Python dict.
load([device])Creates model described by this config and loads weights.
to_dict()Converts this config to a Python dict containing only primitive values.
Attributes
graph_hashn_outputstokenizer_groupsweights_kaggle_idweights_pathmodel_typeinput_sizenum_classes_for_one_hotlayers_sizes- build_model() Module[source]
Creates model described by this config, with randomly initialized weights.
- load(device='cpu') Module[source]
Creates model described by this config and loads weights.
Weights are loaded from weights_path. A config with neither weights_path nor weights_kaggle_id describes an untrained model, and the returned model has randomly initialized weights.
- Parameters:
device – PyTorch device to load the model to.
- Returns:
The model.