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)