(minibatch, optimizer, model, loss_fn)
| 137 | |
| 138 | |
| 139 | def train_step(minibatch, optimizer, model, loss_fn): |
| 140 | node_features = minibatch.node_features["feat"] |
| 141 | labels = minibatch.labels |
| 142 | optimizer.zero_grad() |
| 143 | out = model(minibatch.blocks, node_features) |
| 144 | loss = loss_fn(out, labels) |
| 145 | num_correct = accuracy(out, labels) * labels.size(0) |
| 146 | loss.backward() |
| 147 | optimizer.step() |
| 148 | return loss.detach(), num_correct, labels.size(0) |
| 149 | |
| 150 | |
| 151 | def train_helper( |