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