File size: 2,390 Bytes
1c61c4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
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