Spaces:
Sleeping
Sleeping
File size: 2,571 Bytes
377b913 | 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 67 68 69 70 71 72 | import os
from datetime import datetime
import socket
import torch
import sys
from torch.utils.tensorboard import SummaryWriter
import csv
class Logger(object):
def __init__(self, directory, write, save_freq=10, comment=""):
self.parent_dir = directory
self.write = write
self.dir = ""
self.fig_index = 0
self.model_index = 0
self.save_freq = save_freq
if self.write:
self.boardWriter = SummaryWriter(comment=comment)
self.dir = self.boardWriter.log_dir
self.log(f"Logs from {self.dir}\n{' '.join(sys.argv)}\n")
def write_to_board(self, name, scalars, index=0):
self.log(f"{name} at {index}: {str(scalars)}")
if self.write:
for key, value in scalars.items():
self.boardWriter.add_scalar(f"{name}/{key}", value, index)
def plot_res(self, losses, distances):
if len(losses) == 0 or not self.write:
return
import matplotlib.pyplot as plt
fig, axs = plt.subplots(2)
axs[0].plot(list(range(len(losses))), losses, color='orange')
axs[0].set_xlabel("Steps")
axs[0].set_ylabel("Loss")
axs[0].set_title("Training")
axs[0].set_yscale('log')
for dist in distances:
axs[1].plot(list(range(len(dist))), dist)
axs[1].set_xlabel("Steps")
axs[1].set_ylabel("Distance change")
axs[1].set_title("Training")
if self.fig_index > 0:
os.remove(os.path.join(self.dir, f"res{self.fig_index-1}.png"))
fig.savefig(os.path.join(self.dir, f"res{self.fig_index}.png"))
self.boardWriter.add_figure(f"res{self.fig_index}", fig)
self.fig_index += 1
def log(self, message, step=0):
print(str(message))
if self.write:
# self.boardWriter.add_text("log", str(message), step)
with open(os.path.join(self.dir, "logs.txt"), "a") as logs:
logs.write(str(message) + "\n")
def save_model(self, state_dict, name="dqn.pt", forced=False):
if not self.write:
return
if (forced or
(self.model_index > 0 and self.model_index % self.save_freq == 0)):
torch.save(state_dict, os.path.join(self.dir, name))
def write_locations(self, row):
self.log(str(row))
if self.write:
with open(os.path.join(self.dir, 'results.csv'),
mode='a', newline='') as f:
res_writer = csv.writer(f)
res_writer.writerow(row)
|