CoolFace
Apppublic

parkererickson/LGGM-Text2Graph

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
utils.py137 linesDownload Raw Back to root
1import os2import torch_geometric.utils3from omegaconf import OmegaConf, open_dict4from torch_geometric.utils import to_dense_adj, to_dense_batch5import torch6import omegaconf7import wandb8 9def create_folders(args):10    try:11        # os.makedirs('checkpoints')12        os.makedirs('graphs')13        os.makedirs('chains')14    except OSError:15        pass16 17    try:18        # os.makedirs('checkpoints/' + args.general.name)19        os.makedirs('graphs/' + args.general.name)20        os.makedirs('chains/' + args.general.name)21    except OSError:22        pass23 24 25def normalize(X, E, y, norm_values, norm_biases, node_mask):26    X = (X - norm_biases[0]) / norm_values[0]27    E = (E - norm_biases[1]) / norm_values[1]28    y = (y - norm_biases[2]) / norm_values[2]29 30    diag = torch.eye(E.shape[1], dtype=torch.bool).unsqueeze(0).expand(E.shape[0], -1, -1)31    E[diag] = 032 33    return PlaceHolder(X=X, E=E, y=y).mask(node_mask)34 35 36def unnormalize(X, E, y, norm_values, norm_biases, node_mask, collapse=False):37    """38    X : node features39    E : edge features40    y : global features`41    norm_values : [norm value X, norm value E, norm value y]42    norm_biases : same order43    node_mask44    """45    X = (X * norm_values[0] + norm_biases[0])46    E = (E * norm_values[1] + norm_biases[1])47    y = y * norm_values[2] + norm_biases[2]48 49    return PlaceHolder(X=X, E=E, y=y).mask(node_mask, collapse)50 51 52def to_dense(x, edge_index, edge_attr, batch):53    X, node_mask = to_dense_batch(x=x, batch=batch)54    # node_mask = node_mask.float()55    edge_index, edge_attr = torch_geometric.utils.remove_self_loops(edge_index, edge_attr)56    # TODO: carefully check if setting node_mask as a bool breaks the continuous case57    max_num_nodes = X.size(1)58    E = to_dense_adj(edge_index=edge_index, batch=batch, edge_attr=edge_attr, max_num_nodes=max_num_nodes)59    E = encode_no_edge(E)60 61    return PlaceHolder(X=X, E=E, y=None), node_mask62 63 64def encode_no_edge(E):65    assert len(E.shape) == 466    if E.shape[-1] == 0:67        return E68    no_edge = torch.sum(E, dim=3) == 069    first_elt = E[:, :, :, 0]70    first_elt[no_edge] = 171    E[:, :, :, 0] = first_elt72    diag = torch.eye(E.shape[1], dtype=torch.bool).unsqueeze(0).expand(E.shape[0], -1, -1)73    E[diag] = 074    return E75 76 77def update_config_with_new_keys(cfg, saved_cfg):78    saved_general = saved_cfg.general79    saved_train = saved_cfg.train80    saved_model = saved_cfg.model81 82    for key, val in saved_general.items():83        OmegaConf.set_struct(cfg.general, True)84        with open_dict(cfg.general):85            if key not in cfg.general.keys():86                setattr(cfg.general, key, val)87 88    OmegaConf.set_struct(cfg.train, True)89    with open_dict(cfg.train):90        for key, val in saved_train.items():91            if key not in cfg.train.keys():92                setattr(cfg.train, key, val)93 94    OmegaConf.set_struct(cfg.model, True)95    with open_dict(cfg.model):96        for key, val in saved_model.items():97            if key not in cfg.model.keys():98                setattr(cfg.model, key, val)99    return cfg100 101 102class PlaceHolder:103    def __init__(self, X, E, y):104        self.X = X105        self.E = E106        self.y = y107 108    def type_as(self, x: torch.Tensor):109        """ Changes the device and dtype of X, E, y. """110        self.X = self.X.type_as(x)111        self.E = self.E.type_as(x)112        self.y = self.y.type_as(x)113        return self114 115    def mask(self, node_mask, collapse=False):116        x_mask = node_mask.unsqueeze(-1)          # bs, n, 1117        e_mask1 = x_mask.unsqueeze(2)             # bs, n, 1, 1118        e_mask2 = x_mask.unsqueeze(1)             # bs, 1, n, 1119 120        if collapse:121            self.X = torch.argmax(self.X, dim=-1)122            self.E = torch.argmax(self.E, dim=-1)123 124            self.X[node_mask == 0] = - 1125            self.E[(e_mask1 * e_mask2).squeeze(-1) == 0] = - 1126        else:127            self.X = self.X * x_mask128            self.E = self.E * e_mask1 * e_mask2129            assert torch.allclose(self.E, torch.transpose(self.E, 1, 2))130        return self131    132def setup_wandb(cfg):133    config_dict = omegaconf.OmegaConf.to_container(cfg, resolve=True, throw_on_missing=True)134    kwargs = {'name': cfg.general.name, 'project': f'graph_ddm_{cfg.dataset.name}', 'config': config_dict,135              'settings': wandb.Settings(_disable_stats=True), 'reinit': True, 'mode': cfg.general.wandb}136    wandb.init(**kwargs)137    wandb.save('*.txt')