File size: 1,872 Bytes
4acbfc7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 | import os
import time
from data.dataset import TextDataset
from models.model import WriteViT
from params import *
def main():
init_project()
TextDatasetObj = TextDataset(num_examples = NUM_EXAMPLES)
dataset = torch.utils.data.DataLoader(
TextDatasetObj,
batch_size=batch_size,
shuffle=True,
num_workers=0,
pin_memory=True, drop_last=True,
collate_fn=TextDatasetObj.collate_fn)
model = WriteViT(backbone=BACKBONE).to(DEVICE)
os.makedirs('saved_models', exist_ok = True)
MODEL_PATH = os.path.join('saved_models', EXP_NAME)
if os.path.isdir(MODEL_PATH) and RESUME:
model.load_state_dict(torch.load(MODEL_PATH+'/model.pth'))
print (MODEL_PATH+' : Model loaded Successfully')
else:
if not os.path.isdir(MODEL_PATH): os.mkdir(MODEL_PATH)
for epoch in range(EPOCHS):
start_time = time.time()
for i,data in enumerate(dataset):
if (i % NUM_CRITIC_GOCR_TRAIN) == 0:
model._set_input(data)
model.optimize_G_only()
model.optimize_G_step()
if (i % NUM_CRITIC_DOCR_TRAIN) == 0:
model._set_input(data)
model.optimize_D_OCR_W()
model.optimize_D_OCR_W_step()
end_time = time.time()
losses = model.get_current_losses()
print ({'EPOCH':epoch, 'TIME':end_time-start_time, 'LOSSES': losses})
if epoch % SAVE_MODEL == 0: torch.save(model.state_dict(), MODEL_PATH+ '/model.pth')
if epoch % SAVE_MODEL_HISTORY == 0: torch.save(model.state_dict(), MODEL_PATH+ '/model'+str(epoch)+'.pth')
if __name__ == "__main__":
main()
|