Changes
2 changed files (+9/-6)
-
-
@@ -1,3 +1,5 @@from collections import Counter import torch
-
-
-
@@ -2,6 +2,7 @@ from argparse import ArgumentParserimport torch from torch import nn from torch.utils.data import DataLoader from dataset import Dataset from model import Model
-
@@ -23,7 +24,7 @@ print(len(dataloader))# Prepare model model = Model(dataset, 512, 512, 3, 0.2).to(device) model = Model(dataset, 512, 512, 3, 0.2).to(args.device) print(model)
-
@@ -35,8 +36,8 @@ for t in range(args.epochs):model.train() state_h, state_c = net.zero_state(flags.batch_size) state_h = state_h.to(device) state_c = state_c.to(device) state_h = state_h.to(args.device) state_c = state_c.to(args.device) iteration = 0
-
@@ -45,8 +46,8 @@ for t in range(args.epochs):optimizer.zero_grad() x = torch.tensor(x).to(device) y = torch.tensor(y).to(device) x = torch.tensor(x).to(args.device) y = torch.tensor(y).to(args.device) # Compute prediction error logits, (state_h, state_c) = net(x, (state_h, state_c))
-
@@ -71,7 +72,7 @@ for t in range(args.epochs):'Loss: {}'.format(loss_value)) if iteration % 1000 == 0: predict(device, net, flags.initial_words, n_vocab, predict(args.device, net, flags.initial_words, n_vocab, vocab_to_int, int_to_vocab, top_k=3) torch.save(net.state_dict(), 'checkpoint/model-{}.pth'.format(iteration))
-