import torch import numpy as np import sys def save_embed_npz(dataset, model_path, data_path, model_name): sys.path.append('./' + str(model_path)) experiment_path = model_path / 'experiments' / (dataset + '_' + model_name) / 'best_model.pth' save_path = data_path.parent / 'embed_model' / (dataset + '_' + model_name + '_embed.npz') model = torch.load(experiment_path) if model_name == 'ComplEx': rel_embeddings = torch.sqrt(torch.square(model.rel_re_embeddings.weight.data) + torch.square(model.rel_im_embeddings.weight.data)) ent_embeddings = torch.sqrt(torch.square(model.ent_re_embeddings.weight.data) + torch.square(model.ent_im_embeddings.weight.data)) else: rel_embeddings = model.rel_embeddings.weight.data ent_embeddings = model.ent_embeddings.weight.data eM = np.array(ent_embeddings.cpu()) rM = np.array(rel_embeddings.cpu()) if not save_path.parent.exists(): save_path.parent.mkdir() np.savez(save_path, eM=eM, rM=rM) print(model_name+' saved successfully!!!') if __name__ == '__main__': save_embed_npz('wordnet', 'TransE')