graphist-v2 / modeling /models /__init__.py
Ace3Z's picture
Fix stale paths and resync the shipped modeling code
36d3a9f verified
Raw
History Blame Contribute Delete
2.39 kB
from .edcoder import PreModel
def build_model(args):
num_hidden = args.num_hidden
num_layers = args.num_layers
encoder_type = args.encoder
decoder_type = args.decoder
mask_rate = args.mask_rate
drop_edge_rate = args.drop_edge_rate
replace_rate = args.replace_rate
batchnorm = args.batchnorm
# Encoder between-layer norm: none | layer | graph. Back-compat: the older
# boolean --encoder_layernorm maps to "layer" when --encoder_norm is unset.
encoder_norm = getattr(args, "encoder_norm", "none")
if encoder_norm == "none" and getattr(args, "encoder_layernorm", False):
encoder_norm = "layer"
# Learnable normalization of the projected input embedding (encoder stem).
encoder_input_norm = getattr(args, "input_norm", "none")
# When False, the edge distance (col 0) is used only as the aggregation
# weight, not fed into the edge feature projection.
edge_distance_in_proj = getattr(args, "edge_distance_in_proj", True)
# VICReg-style anti-collapse regularizer weights (0 = off). Applied on the
# graph-level pooled embedding during training. Absent on eval-time args
# (generate_embs) -> getattr defaults keep those code paths unaffected.
vicreg_var_weight = getattr(args, "vicreg_var_weight", 0.0)
vicreg_cov_weight = getattr(args, "vicreg_cov_weight", 0.0)
vicreg_gamma = getattr(args, "vicreg_gamma", 1.0)
activation = args.activation
loss_fn = args.loss_fn
alpha_l = args.alpha_l
concat_hidden = args.concat_hidden
num_features = args.num_features
num_edge_features = args.num_edge_features
model = PreModel(
in_dim=int(num_features),
edge_in_dim=int(num_edge_features),
num_hidden=int(num_hidden),
num_layers=num_layers,
activation=activation,
encoder_type=encoder_type,
decoder_type=decoder_type,
mask_rate=mask_rate,
loss_fn=loss_fn,
drop_edge_rate=drop_edge_rate,
replace_rate=replace_rate,
alpha_l=alpha_l,
concat_hidden=concat_hidden,
batchnorm=batchnorm,
encoder_norm=encoder_norm,
encoder_input_norm=encoder_input_norm,
edge_distance_in_proj=edge_distance_in_proj,
vicreg_var_weight=vicreg_var_weight,
vicreg_cov_weight=vicreg_cov_weight,
vicreg_gamma=vicreg_gamma,
)
return model