diff --git a/fastai_contrib/data.py b/fastai_contrib/data.py index 274d1f6..d7be826 100644 --- a/fastai_contrib/data.py +++ b/fastai_contrib/data.py @@ -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." diff --git a/fastai_contrib/learner.py b/fastai_contrib/learner.py index e060b7a..bdcea7c 100644 --- a/fastai_contrib/learner.py +++ b/fastai_contrib/learner.py @@ -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