| 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') | |