Decode function
(args, s_arc, s_rel, mask)
| 227 | |
| 228 | |
| 229 | def 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 | |
| 246 | def save(path, args, model, optimizer): |
no outgoing calls
no test coverage detected