adding tensorboard writer for training loss and cv loss
This commit is contained in:
@@ -33,6 +33,10 @@ You can then train a default behavior cloning policy with the following. Be sure
|
|||||||
```
|
```
|
||||||
python src/main.py --train
|
python src/main.py --train
|
||||||
```
|
```
|
||||||
|
You can run tensorboard by running the following and opening `localhost:6006` (or alternatively port-forwarding 6006 from the remote server)
|
||||||
|
```
|
||||||
|
tensorboard --logdir output/
|
||||||
|
```
|
||||||
You can then test the learned policy with the following, and see the animation file in `output/`:
|
You can then test the learned policy with the following, and see the animation file in `output/`:
|
||||||
```
|
```
|
||||||
python src/main.py --test
|
python src/main.py --test
|
||||||
|
|||||||
@@ -4,3 +4,4 @@ sklearn
|
|||||||
pytest
|
pytest
|
||||||
json5
|
json5
|
||||||
tqdm
|
tqdm
|
||||||
|
tensorboard
|
||||||
|
|||||||
16
src/bc/bc.py
16
src/bc/bc.py
@@ -2,11 +2,11 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
import pickle
|
import pickle
|
||||||
#from torch.utils.tensorboard import SummaryWriter
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
|
|
||||||
from src.policies import DeepSetsPolicy
|
from src.policies import DeepSetsPolicy
|
||||||
from src.util.transform import MinMaxScaler
|
from src.util.transform import MinMaxScaler
|
||||||
import json5
|
from tqdm import tqdm
|
||||||
|
|
||||||
class BehaviorCloningPolicy():
|
class BehaviorCloningPolicy():
|
||||||
"""
|
"""
|
||||||
@@ -115,9 +115,9 @@ def generate_transforms(dataset):
|
|||||||
def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
|
def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
|
||||||
|
|
||||||
# hyperparams
|
# hyperparams
|
||||||
train_epochs = 100
|
train_epochs = 1000
|
||||||
cv_every = 10
|
cv_every = 10
|
||||||
epoch_every = 1
|
epoch_every = 1000
|
||||||
train_batch_size = 64
|
train_batch_size = 64
|
||||||
cv_batch_size = 256 # doesn't matter
|
cv_batch_size = 256 # doesn't matter
|
||||||
learning_rate = 1e-3
|
learning_rate = 1e-3
|
||||||
@@ -140,7 +140,10 @@ def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
|
|||||||
loss_fn = nn.HuberLoss(reduction='sum')
|
loss_fn = nn.HuberLoss(reduction='sum')
|
||||||
optimizer = torch.optim.Adam(policy.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
optimizer = torch.optim.Adam(policy.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
||||||
|
|
||||||
for i in range(train_epochs):
|
# generate tensorboard writer
|
||||||
|
writer = SummaryWriter(filestr)
|
||||||
|
|
||||||
|
for i in tqdm(range(train_epochs)):
|
||||||
|
|
||||||
epoch_loss = 0
|
epoch_loss = 0
|
||||||
for (batch_idx, batch) in enumerate(training_loader):
|
for (batch_idx, batch) in enumerate(training_loader):
|
||||||
@@ -160,6 +163,8 @@ def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
|
|||||||
epoch_loss += loss.item() / len(train_dataset)
|
epoch_loss += loss.item() / len(train_dataset)
|
||||||
|
|
||||||
# Write epoch loss
|
# Write epoch loss
|
||||||
|
|
||||||
|
writer.add_scalar('training loss',epoch_loss, i)
|
||||||
if i % epoch_every == 0:
|
if i % epoch_every == 0:
|
||||||
print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss))
|
print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss))
|
||||||
|
|
||||||
@@ -171,6 +176,7 @@ def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
|
|||||||
pred_action = policy(batch)
|
pred_action = policy(batch)
|
||||||
loss = loss_fn(pred_action, batch['action'])
|
loss = loss_fn(pred_action, batch['action'])
|
||||||
cv_loss += loss.item() / len(cv_dataset)
|
cv_loss += loss.item() / len(cv_dataset)
|
||||||
|
writer.add_scalar('cv loss', cv_loss, i)
|
||||||
print('Epoch: {}, CV Loss: {}'.format(i, cv_loss))
|
print('Epoch: {}, CV Loss: {}'.format(i, cv_loss))
|
||||||
|
|
||||||
policy.save_model(filestr)
|
policy.save_model(filestr)
|
||||||
|
|||||||
Reference in New Issue
Block a user