Update to newst fastai

This commit is contained in:
Piotr Czapla
2018-12-01 10:57:20 +01:00
parent 80d4d4da29
commit 6d6ebef1ca
2 changed files with 19 additions and 18 deletions
+4 -6
View File
@@ -16,8 +16,9 @@ class LanguageModelLoader(): # copy of the original LanguageModelLoader
max_len:int=25):
self.dataset,self.bs,self.bptt,self.lm_type,self.shuffle = dataset,bs,bptt,lm_type,shuffle
self.first,self.i,self.iter = True,0,0
self.n = len(np.concatenate(dataset.x.items)) // self.bs
self.n = len(np.concatenate(dataset.x.items)) // self.bs if len(dataset.x.items) > 0 else 0
self.max_len,self.num_workers = max_len,0
self.init_kwargs = dict(bs=bs, bptt=bptt, lm_type=lm_type, shuffle=shuffle, max_len=max_len)
def __iter__(self):
if getattr(self.dataset, 'item', None) is not None:
@@ -41,12 +42,9 @@ class LanguageModelLoader(): # copy of the original LanguageModelLoader
def __getattr__(self,k:str)->Any: return getattr(self.dataset, k)
@property
def batch_size(self):
return self.bs
def batch_size(self): return self.bs
@batch_size.setter
def batch_size(self, v):
self.bs = v
def batch_size(self, v): self.bs = v
def batchify(self, data:np.ndarray) -> LongTensor:
"Split the corpus `data` in batches."
+15 -12
View File
@@ -78,18 +78,21 @@ def convert_weights(wgts:Weights, stoi_wgts:Dict[str,int], itos_new:Collection[s
def convert_weights_with_prefix(wgts:Weights, stoi_wgts:Dict[str,int], itos_new:Collection[str], prefix='') -> Weights:
"Convert the model weights to go with a new vocabulary."
dec_bias, enc_wgts = wgts[prefix+'1.decoder.bias'], wgts[prefix+'0.encoder.weight']
bias_m, wgts_m = dec_bias.mean(0), enc_wgts.mean(0)
new_w = enc_wgts.new_zeros((len(itos_new),enc_wgts.size(1))).zero_()
new_b = dec_bias.new_zeros((len(itos_new),)).zero_()
for i,w in enumerate(itos_new):
r = stoi_wgts[w] if w in stoi_wgts else -1
new_w[i] = enc_wgts[r] if r>=0 else wgts_m
new_b[i] = dec_bias[r] if r>=0 else bias_m
wgts[prefix+'0.encoder.weight'] = new_w
wgts[prefix+'0.encoder_dp.emb.weight'] = new_w.clone()
wgts[prefix+'1.decoder.weight'] = new_w.clone()
wgts[prefix+'1.decoder.bias'] = new_b
if 'model' in wgts:
wgts['model'] = convert_weights_with_prefix(wgts['model'], stoi_wgts, itos_new, prefix)
else:
dec_bias, enc_wgts = wgts[prefix+'1.decoder.bias'], wgts[prefix+'0.encoder.weight']
bias_m, wgts_m = dec_bias.mean(0), enc_wgts.mean(0)
new_w = enc_wgts.new_zeros((len(itos_new),enc_wgts.size(1))).zero_()
new_b = dec_bias.new_zeros((len(itos_new),)).zero_()
for i,w in enumerate(itos_new):
r = stoi_wgts[w] if w in stoi_wgts else -1
new_w[i] = enc_wgts[r] if r>=0 else wgts_m
new_b[i] = dec_bias[r] if r>=0 else bias_m
wgts[prefix+'0.encoder.weight'] = new_w
wgts[prefix+'0.encoder_dp.emb.weight'] = new_w.clone()
wgts[prefix+'1.decoder.weight'] = new_w.clone()
wgts[prefix+'1.decoder.bias'] = new_b
return wgts
#endregion