MCPcopy Create free account
hub / github.com/pytorch/tutorials / train

Function train

intermediate_source/char_rnn_generation_tutorial.py:280–298  ·  view source on GitHub ↗
(category_tensor, input_line_tensor, target_line_tensor)

Source from the content-addressed store, hash-verified

278learning_rate = 0.0005
279
280def train(category_tensor, input_line_tensor, target_line_tensor):
281 target_line_tensor.unsqueeze_(-1)
282 hidden = rnn.initHidden()
283
284 rnn.zero_grad()
285
286 loss = torch.Tensor([0]) # you can also just simply use ``loss = 0``
287
288 for i in range(input_line_tensor.size(0)):
289 output, hidden = rnn(category_tensor, input_line_tensor[i], hidden)
290 l = criterion(output, target_line_tensor[i])
291 loss += l
292
293 loss.backward()
294
295 for p in rnn.parameters():
296 p.data.add_(p.grad.data, alpha=-learning_rate)
297
298 return output, loss.item() / input_line_tensor.size(0)
299
300
301######################################################################

Callers 1

Calls 2

initHiddenMethod · 0.80
backwardMethod · 0.45

Tested by

no test coverage detected