mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-08-21 11:16:18 +08:00
283 lines
14 KiB
Python
283 lines
14 KiB
Python
from collections import defaultdict
|
|
from functools import partial
|
|
from pathlib import Path
|
|
import shutil
|
|
import sys
|
|
import time
|
|
from typing import Any, Dict, Optional, Tuple
|
|
|
|
import hydra
|
|
from hydra.utils import instantiate
|
|
from omegaconf import DictConfig, OmegaConf
|
|
import torch
|
|
import torch.nn as nn
|
|
from tqdm import tqdm
|
|
import wandb
|
|
|
|
from agent import Agent
|
|
from collector import Collector
|
|
from envs import SingleProcessEnv, MultiProcessEnv
|
|
from episode import Episode
|
|
from make_reconstructions import make_reconstructions_from_batch
|
|
from models.actor_critic import ActorCritic
|
|
from models.world_model import WorldModel
|
|
from utils import configure_optimizer, EpisodeDirManager, set_seed
|
|
|
|
|
|
class Trainer:
|
|
def __init__(self, cfg: DictConfig) -> None:
|
|
wandb.init(
|
|
config=OmegaConf.to_container(cfg, resolve=True),
|
|
reinit=True,
|
|
resume=True,
|
|
**cfg.wandb
|
|
)
|
|
|
|
if cfg.common.seed is not None:
|
|
set_seed(cfg.common.seed)
|
|
|
|
self.cfg = cfg
|
|
self.start_epoch = 1
|
|
self.device = torch.device(cfg.common.device)
|
|
|
|
self.ckpt_dir = Path('checkpoints')
|
|
self.media_dir = Path('media')
|
|
self.episode_dir = self.media_dir / 'episodes'
|
|
self.reconstructions_dir = self.media_dir / 'reconstructions'
|
|
|
|
if not cfg.common.resume:
|
|
config_dir = Path('config')
|
|
config_path = config_dir / 'trainer.yaml'
|
|
config_dir.mkdir(exist_ok=False, parents=False)
|
|
shutil.copy('.hydra/config.yaml', config_path)
|
|
wandb.save(str(config_path))
|
|
shutil.copytree(src=(Path(hydra.utils.get_original_cwd()) / "src"), dst="./src")
|
|
shutil.copytree(src=(Path(hydra.utils.get_original_cwd()) / "scripts"), dst="./scripts")
|
|
self.ckpt_dir.mkdir(exist_ok=False, parents=False)
|
|
self.media_dir.mkdir(exist_ok=False, parents=False)
|
|
self.episode_dir.mkdir(exist_ok=False, parents=False)
|
|
self.reconstructions_dir.mkdir(exist_ok=False, parents=False)
|
|
|
|
episode_manager_train = EpisodeDirManager(self.episode_dir / 'train', max_num_episodes=cfg.collection.train.num_episodes_to_save)
|
|
episode_manager_test = EpisodeDirManager(self.episode_dir / 'test', max_num_episodes=cfg.collection.test.num_episodes_to_save)
|
|
self.episode_manager_imagination = EpisodeDirManager(self.episode_dir / 'imagination', max_num_episodes=cfg.evaluation.actor_critic.num_episodes_to_save)
|
|
|
|
def create_env(cfg_env, num_envs):
|
|
env_fn = partial(instantiate, config=cfg_env)
|
|
return MultiProcessEnv(env_fn, num_envs, should_wait_num_envs_ratio=1.0) if num_envs > 1 else SingleProcessEnv(env_fn)
|
|
|
|
if self.cfg.training.should:
|
|
train_env = create_env(cfg.env.train, cfg.collection.train.num_envs)
|
|
self.train_dataset = instantiate(cfg.datasets.train)
|
|
self.train_collector = Collector(train_env, self.train_dataset, episode_manager_train)
|
|
|
|
if self.cfg.evaluation.should:
|
|
test_env = create_env(cfg.env.test, cfg.collection.test.num_envs)
|
|
self.test_dataset = instantiate(cfg.datasets.test)
|
|
self.test_collector = Collector(test_env, self.test_dataset, episode_manager_test)
|
|
|
|
assert self.cfg.training.should or self.cfg.evaluation.should
|
|
env = train_env if self.cfg.training.should else test_env
|
|
|
|
tokenizer = instantiate(cfg.tokenizer)
|
|
world_model = WorldModel(obs_vocab_size=tokenizer.vocab_size, act_vocab_size=env.num_actions, config=instantiate(cfg.world_model))
|
|
actor_critic = ActorCritic(**cfg.actor_critic, act_vocab_size=env.num_actions)
|
|
self.agent = Agent(tokenizer, world_model, actor_critic).to(self.device)
|
|
print(f'{sum(p.numel() for p in self.agent.tokenizer.parameters())} parameters in agent.tokenizer')
|
|
print(f'{sum(p.numel() for p in self.agent.world_model.parameters())} parameters in agent.world_model')
|
|
print(f'{sum(p.numel() for p in self.agent.actor_critic.parameters())} parameters in agent.actor_critic')
|
|
|
|
self.optimizer_tokenizer = torch.optim.Adam(self.agent.tokenizer.parameters(), lr=cfg.training.learning_rate)
|
|
self.optimizer_world_model = configure_optimizer(self.agent.world_model, cfg.training.learning_rate, cfg.training.world_model.weight_decay)
|
|
self.optimizer_actor_critic = torch.optim.Adam(self.agent.actor_critic.parameters(), lr=cfg.training.learning_rate)
|
|
|
|
if cfg.initialization.path_to_checkpoint is not None:
|
|
self.agent.load(**cfg.initialization, device=self.device)
|
|
|
|
if cfg.common.resume:
|
|
self.load_checkpoint()
|
|
|
|
def run(self) -> None:
|
|
|
|
for epoch in range(self.start_epoch, 1 + self.cfg.common.epochs):
|
|
|
|
print(f"\nEpoch {epoch} / {self.cfg.common.epochs}\n")
|
|
start_time = time.time()
|
|
to_log = []
|
|
|
|
if self.cfg.training.should:
|
|
if epoch <= self.cfg.collection.train.stop_after_epochs:
|
|
to_log += self.train_collector.collect(self.agent, epoch, **self.cfg.collection.train.config)
|
|
to_log += self.train_agent(epoch)
|
|
|
|
if self.cfg.evaluation.should and (epoch % self.cfg.evaluation.every == 0):
|
|
self.test_dataset.clear()
|
|
to_log += self.test_collector.collect(self.agent, epoch, **self.cfg.collection.test.config)
|
|
to_log += self.eval_agent(epoch)
|
|
|
|
if self.cfg.training.should:
|
|
self.save_checkpoint(epoch, save_agent_only=not self.cfg.common.do_checkpoint)
|
|
|
|
to_log.append({'duration': (time.time() - start_time) / 3600})
|
|
for metrics in to_log:
|
|
wandb.log({'epoch': epoch, **metrics})
|
|
|
|
self.finish()
|
|
|
|
def train_agent(self, epoch: int) -> None:
|
|
self.agent.train()
|
|
self.agent.zero_grad()
|
|
|
|
metrics_tokenizer, metrics_world_model, metrics_actor_critic = {}, {}, {}
|
|
|
|
cfg_tokenizer = self.cfg.training.tokenizer
|
|
cfg_world_model = self.cfg.training.world_model
|
|
cfg_actor_critic = self.cfg.training.actor_critic
|
|
|
|
if epoch > cfg_tokenizer.start_after_epochs:
|
|
metrics_tokenizer = self.train_component(self.agent.tokenizer, self.optimizer_tokenizer, sequence_length=1, sample_from_start=True, **cfg_tokenizer)
|
|
self.agent.tokenizer.eval()
|
|
|
|
if epoch > cfg_world_model.start_after_epochs:
|
|
metrics_world_model = self.train_component(self.agent.world_model, self.optimizer_world_model, sequence_length=self.cfg.common.sequence_length, sample_from_start=True, tokenizer=self.agent.tokenizer, **cfg_world_model)
|
|
self.agent.world_model.eval()
|
|
|
|
if epoch > cfg_actor_critic.start_after_epochs:
|
|
metrics_actor_critic = self.train_component(self.agent.actor_critic, self.optimizer_actor_critic, sequence_length=1 + self.cfg.training.actor_critic.burn_in, sample_from_start=False, tokenizer=self.agent.tokenizer, world_model=self.agent.world_model, **cfg_actor_critic)
|
|
self.agent.actor_critic.eval()
|
|
|
|
return [{'epoch': epoch, **metrics_tokenizer, **metrics_world_model, **metrics_actor_critic}]
|
|
|
|
def train_component(self, component: nn.Module, optimizer: torch.optim.Optimizer, steps_per_epoch: int, batch_num_samples: int, grad_acc_steps: int, max_grad_norm: Optional[float], sequence_length: int, sample_from_start: bool, **kwargs_loss: Any) -> Dict[str, float]:
|
|
loss_total_epoch = 0.0
|
|
intermediate_losses = defaultdict(float)
|
|
|
|
for _ in tqdm(range(steps_per_epoch), desc=f"Training {str(component)}", file=sys.stdout):
|
|
optimizer.zero_grad()
|
|
for _ in range(grad_acc_steps):
|
|
batch = self.train_dataset.sample_batch(batch_num_samples, sequence_length, sample_from_start)
|
|
batch = self._to_device(batch)
|
|
|
|
losses = component.compute_loss(batch, **kwargs_loss) / grad_acc_steps
|
|
loss_total_step = losses.loss_total
|
|
loss_total_step.backward()
|
|
loss_total_epoch += loss_total_step.item() / steps_per_epoch
|
|
|
|
for loss_name, loss_value in losses.intermediate_losses.items():
|
|
intermediate_losses[f"{str(component)}/train/{loss_name}"] += loss_value / steps_per_epoch
|
|
|
|
if max_grad_norm is not None:
|
|
torch.nn.utils.clip_grad_norm_(component.parameters(), max_grad_norm)
|
|
|
|
optimizer.step()
|
|
|
|
metrics = {f'{str(component)}/train/total_loss': loss_total_epoch, **intermediate_losses}
|
|
return metrics
|
|
|
|
@torch.no_grad()
|
|
def eval_agent(self, epoch: int) -> None:
|
|
self.agent.eval()
|
|
|
|
metrics_tokenizer, metrics_world_model = {}, {}
|
|
|
|
cfg_tokenizer = self.cfg.evaluation.tokenizer
|
|
cfg_world_model = self.cfg.evaluation.world_model
|
|
cfg_actor_critic = self.cfg.evaluation.actor_critic
|
|
|
|
if epoch > cfg_tokenizer.start_after_epochs:
|
|
metrics_tokenizer = self.eval_component(self.agent.tokenizer, cfg_tokenizer.batch_num_samples, sequence_length=1)
|
|
|
|
if epoch > cfg_world_model.start_after_epochs:
|
|
metrics_world_model = self.eval_component(self.agent.world_model, cfg_world_model.batch_num_samples, sequence_length=self.cfg.common.sequence_length, tokenizer=self.agent.tokenizer)
|
|
|
|
if epoch > cfg_actor_critic.start_after_epochs:
|
|
self.inspect_imagination(epoch)
|
|
|
|
if cfg_tokenizer.save_reconstructions:
|
|
batch = self._to_device(self.test_dataset.sample_batch(batch_num_samples=3, sequence_length=self.cfg.common.sequence_length))
|
|
make_reconstructions_from_batch(batch, save_dir=self.reconstructions_dir, epoch=epoch, tokenizer=self.agent.tokenizer)
|
|
|
|
return [metrics_tokenizer, metrics_world_model]
|
|
|
|
@torch.no_grad()
|
|
def eval_component(self, component: nn.Module, batch_num_samples: int, sequence_length: int, **kwargs_loss: Any) -> Dict[str, float]:
|
|
loss_total_epoch = 0.0
|
|
intermediate_losses = defaultdict(float)
|
|
|
|
steps = 0
|
|
pbar = tqdm(desc=f"Evaluating {str(component)}", file=sys.stdout)
|
|
for batch in self.test_dataset.traverse(batch_num_samples, sequence_length):
|
|
batch = self._to_device(batch)
|
|
|
|
losses = component.compute_loss(batch, **kwargs_loss)
|
|
loss_total_epoch += losses.loss_total.item()
|
|
|
|
for loss_name, loss_value in losses.intermediate_losses.items():
|
|
intermediate_losses[f"{str(component)}/eval/{loss_name}"] += loss_value
|
|
|
|
steps += 1
|
|
pbar.update(1)
|
|
|
|
intermediate_losses = {k: v / steps for k, v in intermediate_losses.items()}
|
|
metrics = {f'{str(component)}/eval/total_loss': loss_total_epoch / steps, **intermediate_losses}
|
|
return metrics
|
|
|
|
@torch.no_grad()
|
|
def inspect_imagination(self, epoch: int) -> None:
|
|
mode_str = 'imagination'
|
|
batch = self.test_dataset.sample_batch(batch_num_samples=self.episode_manager_imagination.max_num_episodes, sequence_length=1 + self.cfg.training.actor_critic.burn_in, sample_from_start=False)
|
|
outputs = self.agent.actor_critic.imagine(self._to_device(batch), self.agent.tokenizer, self.agent.world_model, horizon=self.cfg.evaluation.actor_critic.horizon, show_pbar=True)
|
|
|
|
to_log = []
|
|
for i, (o, a, r, d) in enumerate(zip(outputs.observations.cpu(), outputs.actions.cpu(), outputs.rewards.cpu(), outputs.ends.long().cpu())): # Make everything (N, T, ...) instead of (T, N, ...)
|
|
episode = Episode(o, a, r, d, torch.ones_like(d))
|
|
episode_id = (epoch - 1 - self.cfg.training.actor_critic.start_after_epochs) * outputs.observations.size(0) + i
|
|
self.episode_manager_imagination.save(episode, episode_id, epoch)
|
|
|
|
metrics_episode = {k: v for k, v in episode.compute_metrics().__dict__.items()}
|
|
metrics_episode['episode_num'] = episode_id
|
|
metrics_episode['action_histogram'] = wandb.Histogram(episode.actions.numpy(), num_bins=self.agent.world_model.act_vocab_size)
|
|
to_log.append({f'{mode_str}/{k}': v for k, v in metrics_episode.items()})
|
|
|
|
return to_log
|
|
|
|
def _save_checkpoint(self, epoch: int, save_agent_only: bool) -> None:
|
|
torch.save(self.agent.state_dict(), self.ckpt_dir / 'last.pt')
|
|
if not save_agent_only:
|
|
torch.save(epoch, self.ckpt_dir / 'epoch.pt')
|
|
torch.save({
|
|
"optimizer_tokenizer": self.optimizer_tokenizer.state_dict(),
|
|
"optimizer_world_model": self.optimizer_world_model.state_dict(),
|
|
"optimizer_actor_critic": self.optimizer_actor_critic.state_dict(),
|
|
}, self.ckpt_dir / 'optimizer.pt')
|
|
ckpt_dataset_dir = self.ckpt_dir / 'dataset'
|
|
ckpt_dataset_dir.mkdir(exist_ok=True, parents=False)
|
|
self.train_dataset.update_disk_checkpoint(ckpt_dataset_dir)
|
|
if self.cfg.evaluation.should:
|
|
torch.save(self.test_dataset.num_seen_episodes, self.ckpt_dir / 'num_seen_episodes_test_dataset.pt')
|
|
|
|
def save_checkpoint(self, epoch: int, save_agent_only: bool) -> None:
|
|
tmp_checkpoint_dir = Path('checkpoints_tmp')
|
|
shutil.copytree(src=self.ckpt_dir, dst=tmp_checkpoint_dir, ignore=shutil.ignore_patterns('dataset'))
|
|
self._save_checkpoint(epoch, save_agent_only)
|
|
shutil.rmtree(tmp_checkpoint_dir)
|
|
|
|
def load_checkpoint(self) -> None:
|
|
assert self.ckpt_dir.is_dir()
|
|
self.start_epoch = torch.load(self.ckpt_dir / 'epoch.pt') + 1
|
|
self.agent.load(self.ckpt_dir / 'last.pt', device=self.device)
|
|
ckpt_opt = torch.load(self.ckpt_dir / 'optimizer.pt', map_location=self.device)
|
|
self.optimizer_tokenizer.load_state_dict(ckpt_opt['optimizer_tokenizer'])
|
|
self.optimizer_world_model.load_state_dict(ckpt_opt['optimizer_world_model'])
|
|
self.optimizer_actor_critic.load_state_dict(ckpt_opt['optimizer_actor_critic'])
|
|
self.train_dataset.load_disk_checkpoint(self.ckpt_dir / 'dataset')
|
|
if self.cfg.evaluation.should:
|
|
self.test_dataset.num_seen_episodes = torch.load(self.ckpt_dir / 'num_seen_episodes_test_dataset.pt')
|
|
print(f'Successfully loaded model, optimizer and {len(self.train_dataset)} episodes from {self.ckpt_dir.absolute()}.')
|
|
|
|
def _to_device(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
|
return {k: batch[k].to(self.device) for k in batch}
|
|
|
|
def finish(self) -> None:
|
|
wandb.finish()
|