Files
iris_bigvae/notebooks/01_debug_models.ipynb

33 KiB

In [1]:
# autoreload import your package
%load_ext autoreload
%autoreload 2

import gym
import numpy as np
import matplotlib.pyplot as plt
%matplotlib inline
plt.style.use('ggplot')

Debug model components

Using trainer? 💩

Hyrda is really annoying

In [2]:

import os
os.environ['WANDB_MODE'] = 'disabled'

import hydra
from hydra import initialize, initialize_config_module, initialize_config_dir, compose
from omegaconf import OmegaConf

from pathlib import Path
from datetime import datetime

from src.trainer import Trainer


class Trainer2(Trainer):
    
    def load_checkpoint(self, *args, **kwargs):
        pass



ts = datetime.now().strftime("%Y-%m-%d/%H-%M-%S")
run_dir = Path(f"..outputs/{ts}").absolute()
run_dir.mkdir(parents=True, exist_ok=True)
abs_config_dir=os.path.abspath("../config")
os.chdir(run_dir)
# with initialize_config_dir(version_base=None, config_dir=abs_config_dir):
with initialize(version_base=None, config_path="../config"):
    cfg = compose(config_name='trainer', overrides=[
        f'hydra.run.dir={run_dir}',
        # f"initialization.path_to_checkpoint={str(path_to_checkpoint.absolute())}",
        'wandb.mode=disabled',
        "env.train.id=CrafterReward-v1",
        "training.tokenizer.start_after_epochs=1",
        "training.world_model.start_after_epochs=1",
        "training.actor_critic.start_after_epochs=1",
        "training.tokenizer.steps_per_epoch=10",
        "training.world_model.steps_per_epoch=10",
        "training.actor_critic.steps_per_epoch=10",
        "common.do_checkpoint=False",
        "common.resume=True",
        "training.world_model.batch_num_samples=4",
        "training.actor_critic.batch_num_samples=4",
        ])
    print(cfg)

    with run_dir:
        Path('media/episodes/train').mkdir(parents=True, exist_ok=True)
        trainer = Trainer2(cfg)
    trainer
/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html
  from .autonotebook import tqdm as notebook_tqdm
Failed to detect the name of this notebook, you can set it manually with the WANDB_NOTEBOOK_NAME environment variable to enable code saving.
{'wandb': {'mode': 'disabled', 'project': 'iris', 'entity': None, 'name': None, 'group': None, 'tags': None, 'notes': None}, 'initialization': {'path_to_checkpoint': None, 'load_tokenizer': False, 'load_world_model': False, 'load_actor_critic': False}, 'common': {'epochs': 600, 'device': 'cuda:0', 'do_checkpoint': False, 'seed': 0, 'sequence_length': '${world_model.max_blocks}', 'resume': True}, 'collection': {'train': {'num_envs': 1, 'stop_after_epochs': 500, 'num_episodes_to_save': 10, 'config': {'epsilon': 0.01, 'should_sample': True, 'temperature': 1.0, 'num_steps': 200, 'burn_in': '${training.actor_critic.burn_in}'}}, 'test': {'num_envs': 8, 'num_episodes_to_save': '${collection.train.num_episodes_to_save}', 'config': {'epsilon': 0.0, 'should_sample': True, 'temperature': 0.5, 'num_episodes': 16, 'burn_in': '${training.actor_critic.burn_in}'}}}, 'training': {'should': True, 'learning_rate': 0.0001, 'tokenizer': {'batch_num_samples': 128, 'grad_acc_steps': 1, 'max_grad_norm': 10.0, 'start_after_epochs': 1, 'steps_per_epoch': 10}, 'world_model': {'batch_num_samples': 4, 'grad_acc_steps': 1, 'max_grad_norm': 10.0, 'weight_decay': 0.01, 'start_after_epochs': 1, 'steps_per_epoch': 10}, 'actor_critic': {'batch_num_samples': 4, 'grad_acc_steps': 1, 'max_grad_norm': 10.0, 'start_after_epochs': 1, 'steps_per_epoch': 10, 'imagine_horizon': '${common.sequence_length}', 'burn_in': 20, 'gamma': 0.995, 'lambda_': 0.95, 'entropy_weight': 0.001}}, 'evaluation': {'should': True, 'every': 5, 'tokenizer': {'batch_num_samples': '${training.tokenizer.batch_num_samples}', 'start_after_epochs': '${training.tokenizer.start_after_epochs}', 'save_reconstructions': True}, 'world_model': {'batch_num_samples': '${training.world_model.batch_num_samples}', 'start_after_epochs': '${training.world_model.start_after_epochs}'}, 'actor_critic': {'num_episodes_to_save': '${training.actor_critic.batch_num_samples}', 'horizon': '${training.actor_critic.imagine_horizon}', 'start_after_epochs': '${training.actor_critic.start_after_epochs}'}}, 'tokenizer': {'_target_': 'src.models.tokenizer.Tokenizer', 'vocab_size': 2048, 'embed_dim': 2048, 'encoder': {'_target_': 'src.models.tokenizer.Encoder', 'config': {'_target_': 'src.models.tokenizer.EncoderDecoderConfig', 'resolution': 64, 'in_channels': 3, 'z_channels': 2048, 'ch': 64, 'ch_mult': [1, 1, 1, 1, 1], 'num_res_blocks': 2, 'attn_resolutions': [8, 16], 'out_ch': 3, 'dropout': 0.0}}, 'decoder': {'_target_': 'src.models.tokenizer.Decoder', 'config': '${..encoder.config}'}}, 'world_model': {'_target_': 'src.models.TransformerConfig', 'max_blocks': 10, 'num_layers': 1, 'num_heads': 1, 'embed_dim': 2048, 'dropout': 0.1, 'model_name': 'PY007/TinyLlama-1.1B-intermediate-step-715k-1.5T', 'rank': 32, 'tokens_per_block': 17}, 'actor_critic': {'use_original_obs': False, 'lstm_dim': 512}, 'env': {'train': {'_target_': 'src.envs.make_env', 'id': 'CrafterReward-v1', 'size': 64, 'max_episode_steps': 20000, 'noop_max': 30, 'frame_skip': 4, 'done_on_life_loss': True, 'clip_reward': False}, 'test': {'_target_': '${..train._target_}', 'id': '${..train.id}', 'size': '${..train.size}', 'max_episode_steps': 108000, 'noop_max': 1, 'frame_skip': '${..train.frame_skip}', 'done_on_life_loss': False, 'clip_reward': False}, 'keymap': 'atari/${.train.id}'}, 'datasets': {'train': {'_target_': 'src.dataset.EpisodesDatasetRamMonitoring', 'max_ram_usage': '30G', 'name': 'train_dataset'}, 'test': {'_target_': 'src.dataset.EpisodesDataset', 'max_num_episodes': None, 'name': 'test_dataset'}}}
Tokenizer : shape of latent is (2048, 4, 4).
/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torchvision/models/_utils.py:208: UserWarning: The parameter 'pretrained' is deprecated since 0.13 and may be removed in the future, please use 'weights' instead.
  warnings.warn(
/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torchvision/models/_utils.py:223: UserWarning: Arguments other than a weight enum or `None` for 'weights' are deprecated since 0.13 and may be removed in the future. The current behavior is equivalent to passing `weights=VGG16_Weights.IMAGENET1K_V1`. You can also use `weights=VGG16_Weights.DEFAULT` to get the most up-to-date weights.
  warnings.warn(msg)
Using pad_token, but it is not set yet.
trainable params: 50,462,720 || all params: 1,150,511,104 || trainable%: 4.386113252149889
None
32314243 parameters in agent.tokenizer
752979973 parameters in agent.world_model
3224626 parameters in agent.actor_critic
In [ ]:

Trainer train_agent

In [3]:
self=trainer
epoch = 52

# get out first exp
self.train_collector.collect(self.agent, epoch, **self.cfg.collection.train.config)
Out [3]:
Experience collection (train_dataset): 100%|██████████| 200/200 [00:03<00:00, 58.00it/s]
[{'train_dataset/episode_length': 183,
  'train_dataset/episode_return': tensor(0.1000),
  'train_dataset/episode_num': 0,
  'train_dataset/action_histogram': <wandb.sdk.data_types.histogram.Histogram at 0x7f97f7373eb0>},
 {'train_dataset/#episodes': 2,
  'train_dataset/#steps': 200,
  'train_dataset/return': 0.100000024}]
In [4]:
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()
In [5]:
import torch
from torchinfo import summary
import torch
from einops import rearrange

Directly benchmark models

In [6]:
tokenizer = self.agent.tokenizer
world_model = self.agent.world_model
actor_critic = self.agent.actor_critic
In [7]:
batch_num_samples = cfg.training.world_model.batch_num_samples
sequence_length = cfg.common.sequence_length
sample_from_start = False
# train_dataset = instantiate(cfg.datasets.train)
batch_num_samples
Out [7]:
4
In [8]:
batch = self.train_dataset.sample_batch(batch_num_samples, sequence_length, sample_from_start)
batch = {k: v.to(self.device) for k, v in batch.items()}
In [9]:
%%time
self.agent.world_model.compute_loss(batch,  tokenizer=self.agent.tokenizer)
Out [9]:
CPU times: user 190 ms, sys: 4.87 ms, total: 194 ms
Wall time: 195 ms
<src.utils.LossWithIntermediateLosses at 0x7f97ed202a30>
In [10]:
%%time
# TODO: why is this so slow?
cfg_actor_critic = self.cfg.training.actor_critic
self.agent.actor_critic.compute_loss(batch,  tokenizer=self.agent.tokenizer, world_model=self.agent.world_model, **cfg_actor_critic)
Out [10]:
CPU times: user 9.35 s, sys: 17.3 ms, total: 9.37 s
Wall time: 9.37 s
<src.utils.LossWithIntermediateLosses at 0x7f97f274beb0>
In [11]:
%%time
# is this the slow part... yes. damn
actor_critic.imagine(batch, tokenizer, world_model, horizon=10);
CPU times: user 9.11 s, sys: 15.5 ms, total: 9.13 s
Wall time: 9.13 s
In [12]:
# # takes 0.1 s, fast
# wm_env = WorldModelEnv(tokenizer, world_model, device)
# wm_env
In [13]:
%%time
# this takes 0.1 seconds and is run 10+ time. So 1 second. Hmm
from src.envs.world_model_env import WorldModelEnv, Categorical
initial_observations = batch['observations']

# get the right obs
wm_env = WorldModelEnv(self.agent.tokenizer, self.agent.world_model, self.device)
obs = wm_env.reset_from_initial_observations(initial_observations[:, -1])
print(obs.shape)


# make sure hidden states are right
self.agent.actor_critic.reset(obs.shape[0])

torch.Size([4, 3, 64, 64])
CPU times: user 105 ms, sys: 207 µs, total: 105 ms
Wall time: 105 ms
In [14]:
import gc
gc.collect()
torch.cuda.empty_cache()
# obs
In [15]:
%%time
# 700us
# fast, executed 10+ times
outputs_ac = actor_critic(obs)
CPU times: user 1.6 ms, sys: 309 µs, total: 1.9 ms
Wall time: 1.72 ms
In [16]:
outputs_ac.logits_actions.shape
# action_token.shape
Out [16]:
torch.Size([4, 1, 17])
In [17]:
# %%timeit
# slow! takes 1s, executed 10+ times this is the culprit, not the lstm. hmm
k=3
horizon = 6
action_token = Categorical(logits=outputs_ac.logits_actions).sample()
obs, reward, done, _ = wm_env.step(action_token, should_predict_next_obs=(k < horizon - 1))
In [18]:
%%timeit
# 62ms
# this is the slow part again. no grad and eval don't hepl
outputs_wm = world_model(action_token, past_keys_values=wm_env.keys_values_wm)
66.5 ms ± 1.53 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)
In [19]:

num_steps=1
prev_steps=0
sequences = world_model.embedder(action_token, num_steps, prev_steps) + world_model.pos_emb(prev_steps + torch.arange(num_steps, device=action_token.device))
In [20]:
%%timeit
# ofc it's the transformer that's slow. I guess we just call it was more than during training
past_keys_values = wm_env.keys_values_wm
x = world_model.transformer(sequences, past_keys_values)
---------------------------------------------------------------------------
AssertionError                            Traceback (most recent call last)
/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/notebooks/01_debug_models.ipynb Cell 24 line 1
----> <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/notebooks/01_debug_models.ipynb#X64sZmlsZQ%3D%3D?line=0'>1</a> get_ipython().run_cell_magic('timeit', '', "# ofc it's the transformer that's slow. I guess we just call it was more than during training\npast_keys_values = wm_env.keys_values_wm\nx = world_model.transformer(sequences, past_keys_values)\n")

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/interactiveshell.py:2515, in InteractiveShell.run_cell_magic(self, magic_name, line, cell)
   2513 with self.builtin_trap:
   2514     args = (magic_arg_s, cell)
-> 2515     result = fn(*args, **kwargs)
   2517 # The code below prevents the output from being displayed
   2518 # when using magics with decorator @output_can_be_silenced
   2519 # when the last Python token in the expression is a ';'.
   2520 if getattr(fn, magic.MAGIC_OUTPUT_CAN_BE_SILENCED, False):

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/magics/execution.py:1189, in ExecutionMagics.timeit(self, line, cell, local_ns)
   1186         if time_number >= 0.2:
   1187             break
-> 1189 all_runs = timer.repeat(repeat, number)
   1190 best = min(all_runs) / number
   1191 worst = max(all_runs) / number

File ~/miniforge3/lib/python3.9/timeit.py:205, in Timer.repeat(self, repeat, number)
    203 r = []
    204 for i in range(repeat):
--> 205     t = self.timeit(number)
    206     r.append(t)
    207 return r

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/magics/execution.py:173, in Timer.timeit(self, number)
    171 gc.disable()
    172 try:
--> 173     timing = self.inner(it, self.timer)
    174 finally:
    175     if gcold:

File <magic-timeit>:3, in inner(_it, _timer)

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/nn/modules/module.py:1518, in Module._wrapped_call_impl(self, *args, **kwargs)
   1516     return self._compiled_call_impl(*args, **kwargs)  # type: ignore[misc]
   1517 else:
-> 1518     return self._call_impl(*args, **kwargs)

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/nn/modules/module.py:1527, in Module._call_impl(self, *args, **kwargs)
   1522 # If we don't have any hooks, we want to skip the rest of the logic in
   1523 # this function, and just call forward.
   1524 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks
   1525         or _global_backward_pre_hooks or _global_backward_hooks
   1526         or _global_forward_hooks or _global_forward_pre_hooks):
-> 1527     return forward_call(*args, **kwargs)
   1529 try:
   1530     result = None

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/transformer.py:69, in Transformer.forward(self, sequences, past_keys_values)
     66     # k_size = (x.shape[0], x.shape[1], x.shape[1], 1)
     67     # v_size = past_keys_values[0]._v_cache._cache.size()
     68     v_size = (k_size[0], k_size[1], x.shape[1], k_size[3])
---> 69     past_keys_values[0].update(torch.rand(v_size), torch.rand(v_size))
     70 return x

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/kv_caching.py:59, in KVCache.update(self, k, v)
     58 def update(self, k: torch.Tensor, v: torch.Tensor):
---> 59     self._k_cache.update(k)
     60     self._v_cache.update(v)

File /media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/kv_caching.py:33, in Cache.update(self, x)
     31 def update(self, x: torch.Tensor) -> None:
     32     assert (x.ndim == self._cache.ndim) and all([x.size(i) == self._cache.size(i) for i in (0, 1, 3)])
---> 33     assert self._size + x.size(2) <= self._cache.shape[2]
     34     self._cache = AssignWithoutInplaceCheck.apply(self._cache, x, 2, self._size, self._size + x.size(2))
     35     self._size += x.size(2)

AssertionError: 
In [ ]:
# past_keys_values = wm_env.keys_values_wm
# x = world_model.transformer(sequences, past_keys_values)
In [ ]:
%%timeit
# ofc it's the transformer that's slow. I guess we just call it was more than during training
past_keys_values = wm_env.keys_values_wm
x = world_model.transformer(sequences, past_keys_values)
In [ ]:
logits_observations = world_model.head_observations(x, num_steps=num_steps, prev_steps=prev_steps)
logits_rewards = world_model.head_rewards(x, num_steps=num_steps, prev_steps=prev_steps)
logits_ends = world_model.head_ends(x, num_steps=num_steps, prev_steps=prev_steps)

Torchinfo model sizes

In [ ]:
observations = self.agent.tokenizer.preprocess_input(rearrange(batch['observations'], 'b t c h w -> (b t) c h w'))
# z, z_quantized, reconstructions = self.agent.tokenizer(observations, should_preprocess=False, should_postprocess=False)
summary(self.agent.tokenizer, input_data=observations)
In [ ]:


with torch.no_grad():
    obs_tokens = self.agent.tokenizer.encode(batch['observations'], should_preprocess=True).tokens  # (BL, K)

act_tokens = rearrange(batch['actions'], 'b l -> b l 1')
tokens = rearrange(torch.cat((obs_tokens, act_tokens), dim=2), 'b l k1 -> b (l k1)')  # 

summary(self.agent.world_model, input_data=tokens)
In [ ]:
from src.envs.world_model_env import WorldModelEnv
initial_observations = batch['observations']

# get the right obs
wm_env = WorldModelEnv(self.agent.tokenizer, self.agent.world_model, self.device)
obs = wm_env.reset_from_initial_observations(initial_observations[:, -1])
obs.shape


# make sure hidden states are right
self.agent.actor_critic.reset(obs.shape[0])
In [ ]:


from torchinfo import summary
summary(self.agent.actor_critic, input_data=obs)

Debug env

In [ ]:
import minihack
env = gym.make("MiniHack-River-v0", observation_keys=("pixel_crop", "pixel", 'blstats', 'message'))
env.reset() # each reset generates a new environment instance
obs, reward, end, info = env.step(1)  # move agent '@' north
print(obs['pixel_crop'].shape)
plt.imshow(obs['pixel_crop'])
plt.show()

print(obs['pixel'].shape)
plt.imshow(obs['pixel'])

In [ ]:
# # plt.imshow(obs['glyphs_crop'])
# obs['glyphs_crop'].shape
# obs['blstats']
In [ ]:
import minihack
env = gym.make("MiniHack-Room-5x5-v0", observation_keys=("pixel_crop", "pixel", 'blstats', 'message'))
env.reset() # each reset generates a new environment instance
obs, reward, end, info = env.step(1)  # move agent '@' north
print(obs['pixel_crop'].shape)
plt.imshow(obs['pixel_crop'])
plt.show()

print(obs['pixel'].shape)
plt.imshow(obs['pixel'])

In [ ]:
In [ ]:
import minihack
import crafter
env = gym.make("CrafterReward-v1")
env.reset() # each reset generates a new environment instance
obs, reward, end, info = env.step(1)  # move agent '@' north
print(obs.shape)
plt.imshow(obs)
plt.show()
In [ ]: