NivkhGloss / model_loaders /model_loaders.py
lunadunkel's picture
Update model_loaders.py
7a77441
Raw
History Blame Contribute Delete
1.73 kB
import torch
import re
from collections import deque
import numpy as np
from models_code.SegmentationModel import SegmentationModel
from models_code.PosModel import PosModel
from models_code.GlossModel import GlossModel
from dicts.symb_vocab import symb_vocab
from dicts.label_dict import label_dict
from dicts.word_vocab import word_vocab
from dicts.char_vocab import char_vocab
from dicts.pos_label_vocab import pos_label_vocab
from dicts.morpheme_vocab import morpheme_vocab
from dicts.gloss_vocab import gloss_vocab
def load_segm_model(path, device='cpu'):
model = SegmentationModel(vocab_size=len(symb_vocab),
labels_number=len(label_dict), hidden_dim=512, n_layers=3, dropout=0.4,
device=device, window=(3, 6), bpe_vocab_size=2500, use_attention=True,
use_lstm=True, use_bpe=True)
model.load_state_dict(torch.load(path, map_location=device)['model_state_dict'])
model.to(device)
model.eval()
return model
def load_pos_model(path, device='cpu'):
model = PosModel(word_embedding_dim=64, char_embedding_dim=32,
hidden_dim=128, vocab_size=len(word_vocab),
char_vocab_size=len(char_vocab), labels_number=len(pos_label_vocab),
device=device, use_char_ids=True, dropout=0.2)
model.load_state_dict(torch.load(path, map_location=device)['model_state_dict'])
model.to(device)
model.eval()
return model
def load_gloss_model(path, device='cpu'):
model = GlossModel(len(morpheme_vocab), embed_dim=128,
dropout=0.5, bidirectional=False,
num_layers=2, hidden_dim=256,
output_dim=len(gloss_vocab), device=device)
model.load_state_dict(torch.load(path, map_location=device)['model_state_dict'])
model.to(device)
model.eval()
return model