MCPcopy Create free account
hub / github.com/baidu/DDParser / decode

Function decode

ddparser/parser/model.py:229–243  ·  view source on GitHub ↗

Decode function

(args, s_arc, s_rel, mask)

Source from the content-addressed store, hash-verified

227
228
229def decode(args, s_arc, s_rel, mask):
230 """Decode function"""
231 mask = mask.numpy()
232 lens = np.sum(mask, -1)
233 # prevent self-loops
234 arc_preds = layers.argmax(s_arc, -1).numpy()
235 bad = [not utils.istree(seq[:i + 1]) for i, seq in zip(lens, arc_preds)]
236 if args.tree and any(bad):
237 arc_preds[bad] = utils.eisner(s_arc.numpy()[bad], mask[bad])
238 arc_preds = dygraph.to_variable(arc_preds, zero_copy=False)
239 rel_preds = layers.argmax(s_rel, axis=-1)
240 # batch_size, seq_len, _ = rel_preds.shape
241 rel_preds = nn.index_sample(rel_preds, layers.unsqueeze(arc_preds, -1))
242 rel_preds = layers.squeeze(rel_preds, axes=[-1])
243 return arc_preds, rel_preds
244
245
246def save(path, args, model, optimizer):

Callers 2

epoch_evaluateFunction · 0.85
epoch_predictFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected