Loss function
(s_arc, s_rel, arcs, rels, mask)
| 212 | |
| 213 | |
| 214 | def loss_function(s_arc, s_rel, arcs, rels, mask): |
| 215 | """Loss function""" |
| 216 | arcs = nn.masked_select(arcs, mask) |
| 217 | rels = nn.masked_select(rels, mask) |
| 218 | s_arc = nn.masked_select(s_arc, mask) |
| 219 | s_rel = nn.masked_select(s_rel, mask) |
| 220 | s_rel = nn.index_sample(s_rel, layers.unsqueeze(arcs, 1)) |
| 221 | arc_loss = layers.cross_entropy(layers.softmax(s_arc), arcs) |
| 222 | rel_loss = layers.cross_entropy(layers.softmax(s_rel), rels) |
| 223 | |
| 224 | loss = layers.reduce_mean(arc_loss + rel_loss) |
| 225 | |
| 226 | return loss |
| 227 | |
| 228 | |
| 229 | def decode(args, s_arc, s_rel, mask): |
no outgoing calls
no test coverage detected