clab / clab/dynet

Why does dynet update parameters that are not in the computation graph?

Open
#741 5 comments 0 reactions 0 assignees View on GitHub
moderate bug
Dominant language
C++
Stars
3.4k
Forks
701
PR merge metrics
No merged PRs in 30d

Description

I tried to implement the [Stack-propagation: Improved Representation Learning for Syntax](https://arxiv.org/pdf/1603.06598.pdf) paper for POS and CHUNK joint learning using stack-propagation. The code is given below. I have a shared BI-LSTM for both POS and CHUNK, while their MLPs are separate. W1, B1, W2, B2 is the MLP for POS and W1C, B1C, W2C, B2C is the MLP for CHUNK. I optimize the objective stochastically by alternating between two updates. For first iteration I update the POS network with backpropagation. For second iteration I compute both tagger and chunker activations, backpropagate the chunking loss through the stacked architecture to update both CHUNk and POS tagger, ignoring the POS tagger’s softmax layer parameters. I keep doing this for some n epochs.

My question is when I update the tagger network, the chunker parameters (W1C, B1C, ..) also get updated, which should not happen since for that particular iteration these parameters are not even in the computation graph. I also tried `update` parameter to `false` and `nobackprop` operatos but without any success. Please help.

from __future__ import unicode_literals

import io
import re
import sys
import math
import string
import random
import pickle
import enchant
from itertools import count, chain
from argparse import ArgumentParser
from collections import Counter, defaultdict

import dynet as dy
import numpy as np
from gensim.models.word2vec import Word2Vec

class Meta:
def __init__(self):
self.c_dim = 32
self.n_hidden = 64
self.lstm_char_dim = 32
self.lstm_word_dim = 64

class POSTagger():
def __init__(self, model=None, meta=None):
self.domains = set(open('DOMAINS').read().split())
self.isurl = re.compile(r'[a-z][a-z][.][a-z][a-z]').search
self.upunkt = re.compile(r'[.,\\!@#$%^&\'*()_+={\[}\]|";:<>?`~/]')
self.punct_table = dict((ord(char), None) for char in string.punctuation)

self.model = dy.Model()
if model:
self.meta = pickle.load(open('%s.meta' %model, 'rb'))
else:
self.meta = meta
self.WORDS_LOOKUP = self.model.add_lookup_parameters((self.meta.n_words, self.meta.w_dim))
if not model:
for word, V in wvm.vocab.iteritems():
self.WORDS_LOOKUP.init_row(V.index+1, wvm.syn0[V.index])

self.CHARS_LOOKUP = self.model.add_lookup_parameters((self.meta.n_chars, self.meta.c_dim))

# MLP on top of biLSTM outputs 100 -> 32 -> ntags
self.W1 = self.model.add_parameters((self.meta.n_hidden, self.meta.lstm_word_dim*2))
self.W1C = self.model.add_parameters((self.meta.n_hidden, self.meta.lstm_word_dim*2+self.meta.n_hidden))
self.W2 = self.model.add_parameters((self.meta.n_tags, self.meta.n_hidden))
self.W2C = self.model.add_parameters((self.meta.n_ctags, self.meta.n_hidden))
self.B1 = self.model.add_parameters(self.meta.n_hidden)
self.B1C = self.model.add_parameters(self.meta.n_hidden)
self.B2 = self.model.add_parameters(self.meta.n_tags)
self.B2C = self.model.add_parameters(self.meta.n_ctags)

# word-level LSTMs
self.fwdRNN = dy.LSTMBuilder(1, self.meta.w_dim+self.meta.lstm_char_dim*2, self.meta.lstm_word_dim, self.model)
self.bwdRNN = dy.LSTMBuilder(1, self.meta.w_dim+self.meta.lstm_char_dim*2, self.meta.lstm_word_dim, self.model)

# char-level LSTMs
self.cFwdRNN = dy.LSTMBuilder(1, self.meta.c_dim, self.meta.lstm_char_dim, self.model)
self.cBwdRNN = dy.LSTMBuilder(1, self.meta.c_dim, self.meta.lstm_char_dim, self.model)
if model:
self.model.load('%s.dy' %model)

def get_index(self, ow):
idx = None
is_url = False
try:
return self.meta.w2i[ow]
if args.lang == 'eng':
return self.meta.w2i[ow.lower()]
except KeyError:
pass
if ow.translate(self.punct_table).isdigit():
return self.meta.w2i['N-U-M']
if (ow.startswith('http://') or ow.startswith('https://') or ow.startswith('www.')):
is_url = True
elif self.isurl(ow):
tokens = self.upunkt.split(ow)
is_url = any(tk in self.domains for tk in tokens[1:])
if is_url:
idx = self.meta.w2i['U-R-L']
return idx

def word_rep(self, w, test=True):
if not test and random.random() < 0.25:
return self.WORDS_LOOKUP[0]
w_idx = self.get_index(w)
if w_idx is not None:
return self.WORDS_LOOKUP[w_idx]
else:
return self.WORDS_LOOKUP[0]

def char_rep(self, w, cf_init, cb_init):
pad_char = self.meta.c2i["<*>"]
char_ids = [pad_char] + [self.meta.c2i[c] if self.meta.cc[c]>5 else self.meta.c2i['_UNK_'] for c in w] + [pad_char]
char_embs = [self.CHARS_LOOKUP[cid] for cid in char_ids]
fw_exps = cf_init.transduce(char_embs)
bw_exps = cb_init.transduce(reversed(char_embs))
return dy.concatenate([ fw_exps[-1], bw_exps[-1] ])

def build_tagging_graph(self, words, test=False, tag='pos'):
dy.renew_cg()
# parameters -> expressions
w1 = dy.parameter(self.W1)
b1 = dy.parameter(self.B1)
if tag == 'pos':
w2 = dy.parameter(self.W2)
b2 = dy.parameter(self.B2)
else:
w1c = dy.parameter(self.W1C)
b1c = dy.parameter(self.B1C)
w2c = dy.parameter(self.W2C)
b2c = dy.parameter(self.B2C)
if not test:
w1c = dy.dropout(w1c, 0.25)
b1c = dy.dropout(b1c, 0.25)
if test:
self.fwdRNN.disable_dropout()
self.bwdRNN.disable_dropout()
self.cFwdRNN.disable_dropout()
self.cBwdRNN.disable_dropout()
else:
w1 = dy.dropout(w1, 0.25)
b1 = dy.dropout(b1, 0.25)
self.fwdRNN.set_dropout(0.25)
self.bwdRNN.set_dropout(0.25)
self.cFwdRNN.set_dropout(0.25)
self.cBwdRNN.set_dropout(0.25)

# initialize the RNNs
f_init = self.fwdRNN.initial_state()
b_init = self.bwdRNN.initial_state()

cf_init = self.cFwdRNN.initial_state()
cb_init = self.cBwdRNN.initial_state()

# get the word vectors. word_rep(...) returns a 128-dim vector expression for each word.
wembs = [self.word_rep(w, test) for w in words]
cembs = [self.char_rep(w, cf_init, cb_init) for w in words]
wembs = [dy.concatenate([w,c]) for w,c in zip(wembs, cembs)]

# feed word vectors into biLSTM
fw_exps = f_init.transduce(wembs)
bw_exps = b_init.transduce(reversed(wembs))

# biLSTM states
bi_exps = [dy.concatenate([f,b]) for f,b in zip(fw_exps, reversed(bw_exps))]

# feed each biLSTM state to an MLP
exps = []
for x in bi_exps:
if tag == 'pos':
r_p = w2*(dy.rectify(w1 * x) + b1) + b2
exps.append(r_p)
continue
r_c = w2c*(dy.rectify(w1c * dy.concatenate([w1*x+b1, x])) + b1c) + b2c
exps.append(r_c)

return exps

def sent_loss(self, words, tags, ctags, tag='pos'):
vecs = self.build_tagging_graph(words, tag=tag)
errs = []
for v,t,ct in zip(vecs,tags,ctags):
if tag == 'pos':
tid = self.meta.t2i[t]
else:
tid = self.meta.ct2i[ct]
err = dy.pickneglogsoftmax(v, tid)
errs.append(err)
return dy.esum(errs)

def tag_sent(self, words, tag='pos'):
p_vecs = self.build_tagging_graph(words, test=True, tag=tag)
p_vecs = [dy.softmax(v) for v in p_vecs]
p_probs = [v.npvalue() for v in p_vecs]
tags, ctags = [], []
for p_prb in p_probs:
xtag = np.argmax(p_prb)
if tag == 'pos':
tags.append(self.meta.i2t[xtag])
else:
tags.append(self.meta.i2ct[xtag])
return zip(words, tags)

def read(fname):
data = []
sent = []
pid = 3 if args.ud else 4
fp = io.open(fname, encoding='utf-8')
for i,line in enumerate(fp):
line = line.split()
if not line:
data.append(sent)
sent = []
else:
w,p,c = line[1], line[pid], line[5]
sent.append((w,p,c))
if sent: data.append(sent)
return data

def eval(dev):
good_sent = bad_sent = good = bad = 0.0
gall, pall = [], []
cgood_sent = cbad_sent = cgood = cbad = 0.0
cgall, cpall = [], []
for sent in dev:
words, gpos, gchunk = zip(*sent)
#gall.extend(golds)
words, ppos = zip(*tagger.tag_sent(words, tag='pos'))
words, pchunk = zip(*tagger.tag_sent(words, tag='chunk'))
#pall.extend(tags)
if list(ppos) == list(gpos): good_sent += 1
else: bad_sent += 1
if list(pchunk) == list(gchunk): cgood_sent += 1
else: cbad_sent += 1
for go,gu in zip(gpos,ppos):
if go == gu: good += 1
else: bad += 1
for go,gu in zip(gchunk,pchunk):
if go == gu: cgood += 1
else: cbad += 1
#print(cr(gall, pall, digits=4))
print(good/(good+bad), good_sent/(good_sent+bad_sent))
print(cgood/(cgood+cbad), cgood_sent/(cgood_sent+cbad_sent))
return good/(good+bad)

def train_tagger(train):
pr_acc = 0.0
rate_decay = 0.25
num_tagged, cum_loss = 0, 0
for ITER in xrange(args.iter):
save = False
random.shuffle(train)
#print(dy.parameter(tagger.W2).npvalue())
for i,s in enumerate(train,1):
if i > 0 and i % 1000 == 0: # print status
trainer.status()
print(cum_loss / num_tagged)
cum_loss, num_tagged = 0, 0
if i == len(train): # eval on dev
new_acc = eval(dev)
if new_acc > pr_acc:
pr_acc = new_acc
save = True
# train on sent
words, gpos, gchunk = zip(*s)
if ITER%2==0:# or ITER < 5:
loss_exp = tagger.sent_loss(words, gpos, gchunk, tag='pos')
else:
loss_exp = tagger.sent_loss(words, gpos, gchunk, tag='chunk')
cum_loss += loss_exp.scalar_value()
num_tagged += len(gpos)
loss_exp.backward()
trainer.update()
print("epoch %r finished" % ITER)
trainer.update()
if save:
save = False
print('Save Point:: %d' %ITER)
if args.save_model:
tagger.model.save('%s.dy' %args.save_model)
sys.stdout.flush()

if __name__ == '__main__':
parser = ArgumentParser(description="POS Tagger")
group = parser.add_mutually_exclusive_group()
parser.add_argument('--dynet-mem')
parser.add_argument('--dynet-gpu')
parser.add_argument('--dynet-seed', dest='seed', type=int)
parser.add_argument('--train')
parser.add_argument('--dev')
parser.add_argument('--embd')
parser.add_argument('--lang')
parser.add_argument('--trainer')
parser.add_argument('--ud', type=int)
parser.add_argument('--iter', type=int, default=500)
parser.add_argument('--evec', type=int)
group.add_argument('--save-model', dest='save_model')
group.add_argument('--load-model', dest='load_model')
args = parser.parse_args()
np.random.seed(args.seed)
random.seed(args.seed)

meta = Meta()
if args.dev:
dev = read(args.dev)
if not args.load_model:
wvm = Word2Vec.load_word2vec_format(args.embd, binary=args.evec)
meta.w_dim = wvm.syn0.shape[1]
meta.n_words = wvm.syn0.shape[0]+1

train = read(args.train)
dev = read(args.dev)
tags, chars, ctags = set(), set(), set()
meta.cc = Counter()
for sent in train:
for w,p,c in sent:
tags.add(p)
ctags.add(c)
chars.update(w)
meta.cc.update(w)
chars.update(['_UNK_', '<*>'])
meta.n_tags = len(tags)
meta.n_ctags = len(ctags)
meta.n_chars = len(chars)
meta.i2t = dict(enumerate(tags))
meta.t2i = {t:i for i,t in meta.i2t.items()}
meta.i2ct = dict(enumerate(ctags))
meta.ct2i = {t:i for i,t in meta.i2ct.items()}
meta.c2i = dict(zip(chars, range(meta.n_chars)))

meta.w2i = {}
for w in wvm.vocab:
meta.w2i[w] = wvm.vocab[w].index + 1

if args.save_model:
pickle.dump(meta, open('%s.meta' %args.save_model, 'wb'))
if args.load_model:
tagger = POSTagger(model=args.load_model)
eval(dev)
else:
tagger = POSTagger(meta=meta)
trainers = {
'momsgd' : dy.MomentumSGDTrainer(tagger.model),#, edecay=0.25),
'adam' : dy.AdamTrainer(tagger.model),# edecay=0.25),
'simsgd' : dy.SimpleSGDTrainer(tagger.model),#, edecay=0.25),
'adagrad' : dy.AdagradTrainer(tagger.model),#, edecay=0.25),
'adadelta' : dy.AdadeltaTrainer(tagger.model)#, edecay=0.25)
}
trainer = trainers[args.trainer]
train_tagger(train)

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.