mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-10 12:13:25 +08:00
debugging model size and speed
This commit is contained in:
Vendored
+26
-1
@@ -4,6 +4,28 @@
|
||||
// For more information, visit: https://go.microsoft.com/fwlink/?linkid=830387
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "test",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${workspaceFolder}/src/main.py",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false,
|
||||
"autoReload": {"enable": true,},
|
||||
"env": {"WANDB_MODE":"disabled"},
|
||||
"args": [
|
||||
"'wandb.mode=disabled",
|
||||
// "env.train.id=BreakoutNoFrameskip-v4",
|
||||
"env.train.id=CrafterReward-v1",
|
||||
// # make it start early
|
||||
"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",
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "main",
|
||||
"type": "python",
|
||||
@@ -14,11 +36,14 @@
|
||||
"autoReload": {"enable": true,},
|
||||
"env": {"WANDB_MODE":"disabled"},
|
||||
"args": [
|
||||
"env.train.id=BreakoutNoFrameskip-v4",
|
||||
"'wandb.mode=disabled",
|
||||
// "env.train.id=BreakoutNoFrameskip-v4",
|
||||
"env.train.id=CrafterReward-v1",
|
||||
// # make it start early
|
||||
"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",
|
||||
|
||||
@@ -1 +1,2 @@
|
||||
use_original_obs: False
|
||||
lstm_dim: 512
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
train:
|
||||
_target_: dataset.EpisodesDatasetRamMonitoring
|
||||
_target_: src.dataset.EpisodesDatasetRamMonitoring
|
||||
max_ram_usage: 30G
|
||||
name: train_dataset
|
||||
test:
|
||||
_target_: dataset.EpisodesDataset
|
||||
_target_: src.dataset.EpisodesDataset
|
||||
max_num_episodes: null
|
||||
name: test_dataset
|
||||
|
||||
Vendored
+2
-2
@@ -1,5 +1,5 @@
|
||||
train:
|
||||
_target_: envs.make_atari
|
||||
_target_: src.envs.make_env
|
||||
id: null
|
||||
size: 64
|
||||
max_episode_steps: 20000
|
||||
@@ -18,4 +18,4 @@ test:
|
||||
done_on_life_loss: False
|
||||
clip_reward: False
|
||||
|
||||
keymap: atari/${.train.id}
|
||||
keymap: atari/${.train.id}
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
_target_: models.tokenizer.Tokenizer
|
||||
_target_: src.models.tokenizer.Tokenizer
|
||||
|
||||
vocab_size: 2048
|
||||
embed_dim: 2048
|
||||
encoder:
|
||||
_target_: models.tokenizer.Encoder
|
||||
_target_: src.models.tokenizer.Encoder
|
||||
config:
|
||||
_target_: models.tokenizer.EncoderDecoderConfig
|
||||
_target_: src.models.tokenizer.EncoderDecoderConfig
|
||||
resolution: 64
|
||||
in_channels: 3
|
||||
z_channels: 2048
|
||||
@@ -16,5 +16,5 @@ encoder:
|
||||
out_ch: 3
|
||||
dropout: 0.0
|
||||
decoder:
|
||||
_target_: models.tokenizer.Decoder
|
||||
_target_: src.models.tokenizer.Decoder
|
||||
config: ${..encoder.config}
|
||||
|
||||
+5
-1
@@ -67,7 +67,7 @@ training:
|
||||
start_after_epochs: 25
|
||||
steps_per_epoch: 200
|
||||
actor_critic:
|
||||
batch_num_samples: 32
|
||||
batch_num_samples: 16
|
||||
grad_acc_steps: 1
|
||||
max_grad_norm: 10.0
|
||||
start_after_epochs: 50
|
||||
@@ -92,3 +92,7 @@ evaluation:
|
||||
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}
|
||||
|
||||
hydra:
|
||||
job:
|
||||
chdir: True
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
_target_: models.TransformerConfig
|
||||
_target_: src.models.TransformerConfig
|
||||
max_blocks: 10 # this is the rollout length when training policy
|
||||
num_layers: 1
|
||||
num_heads: 1
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
|
||||
breakout:
|
||||
python src/main.py env.train.id=BreakoutNoFrameskip-v4
|
||||
|
||||
crafter:
|
||||
python src/main.py env.train.id=CrafterReward-v1
|
||||
|
||||
minihack:
|
||||
python src/main.py env.train.id=MiniHack-River-v0
|
||||
|
||||
@@ -0,0 +1,747 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# autoreload import your package\n",
|
||||
"%load_ext autoreload\n",
|
||||
"%autoreload 2\n",
|
||||
"\n",
|
||||
"import gym\n",
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"%matplotlib inline\n",
|
||||
"plt.style.use('ggplot')\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Debug model components\n",
|
||||
"\n",
|
||||
"### Using trainer? :poop:\n",
|
||||
"\n",
|
||||
"Hyrda is really annoying\n",
|
||||
"\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/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\n",
|
||||
" from .autonotebook import tqdm as notebook_tqdm\n",
|
||||
"Failed to detect the name of this notebook, you can set it manually with the WANDB_NOTEBOOK_NAME environment variable to enable code saving.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{'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'}}}\n",
|
||||
"Tokenizer : shape of latent is (2048, 4, 4).\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/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.\n",
|
||||
" warnings.warn(\n",
|
||||
"/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.\n",
|
||||
" warnings.warn(msg)\n",
|
||||
"Using pad_token, but it is not set yet.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"trainable params: 50,462,720 || all params: 1,150,511,104 || trainable%: 4.386113252149889\n",
|
||||
"None\n",
|
||||
"32314243 parameters in agent.tokenizer\n",
|
||||
"752979973 parameters in agent.world_model\n",
|
||||
"3224626 parameters in agent.actor_critic\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"import os\n",
|
||||
"os.environ['WANDB_MODE'] = 'disabled'\n",
|
||||
"\n",
|
||||
"import hydra\n",
|
||||
"from hydra import initialize, initialize_config_module, initialize_config_dir, compose\n",
|
||||
"from omegaconf import OmegaConf\n",
|
||||
"\n",
|
||||
"from pathlib import Path\n",
|
||||
"from datetime import datetime\n",
|
||||
"\n",
|
||||
"from src.trainer import Trainer\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"class Trainer2(Trainer):\n",
|
||||
" \n",
|
||||
" def load_checkpoint(self, *args, **kwargs):\n",
|
||||
" pass\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"ts = datetime.now().strftime(\"%Y-%m-%d/%H-%M-%S\")\n",
|
||||
"run_dir = Path(f\"..outputs/{ts}\").absolute()\n",
|
||||
"run_dir.mkdir(parents=True, exist_ok=True)\n",
|
||||
"abs_config_dir=os.path.abspath(\"../config\")\n",
|
||||
"os.chdir(run_dir)\n",
|
||||
"# with initialize_config_dir(version_base=None, config_dir=abs_config_dir):\n",
|
||||
"with initialize(version_base=None, config_path=\"../config\"):\n",
|
||||
" cfg = compose(config_name='trainer', overrides=[\n",
|
||||
" f'hydra.run.dir={run_dir}',\n",
|
||||
" # f\"initialization.path_to_checkpoint={str(path_to_checkpoint.absolute())}\",\n",
|
||||
" 'wandb.mode=disabled',\n",
|
||||
" \"env.train.id=CrafterReward-v1\",\n",
|
||||
" \"training.tokenizer.start_after_epochs=1\",\n",
|
||||
" \"training.world_model.start_after_epochs=1\",\n",
|
||||
" \"training.actor_critic.start_after_epochs=1\",\n",
|
||||
" \"training.tokenizer.steps_per_epoch=10\",\n",
|
||||
" \"training.world_model.steps_per_epoch=10\",\n",
|
||||
" \"training.actor_critic.steps_per_epoch=10\",\n",
|
||||
" \"common.do_checkpoint=False\",\n",
|
||||
" \"common.resume=True\",\n",
|
||||
" \"training.world_model.batch_num_samples=4\",\n",
|
||||
" \"training.actor_critic.batch_num_samples=4\",\n",
|
||||
" ])\n",
|
||||
" print(cfg)\n",
|
||||
"\n",
|
||||
" with run_dir:\n",
|
||||
" Path('media/episodes/train').mkdir(parents=True, exist_ok=True)\n",
|
||||
" trainer = Trainer2(cfg)\n",
|
||||
" trainer\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"\n",
|
||||
"## Trainer train_agent\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Experience collection (train_dataset): 100%|██████████| 200/200 [00:03<00:00, 58.00it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"[{'train_dataset/episode_length': 183,\n",
|
||||
" 'train_dataset/episode_return': tensor(0.1000),\n",
|
||||
" 'train_dataset/episode_num': 0,\n",
|
||||
" 'train_dataset/action_histogram': <wandb.sdk.data_types.histogram.Histogram at 0x7f97f7373eb0>},\n",
|
||||
" {'train_dataset/#episodes': 2,\n",
|
||||
" 'train_dataset/#steps': 200,\n",
|
||||
" 'train_dataset/return': 0.100000024}]"
|
||||
]
|
||||
},
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"self=trainer\n",
|
||||
"epoch = 52\n",
|
||||
"\n",
|
||||
"# get out first exp\n",
|
||||
"self.train_collector.collect(self.agent, epoch, **self.cfg.collection.train.config)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"self.agent.train()\n",
|
||||
"self.agent.zero_grad()\n",
|
||||
"\n",
|
||||
"metrics_tokenizer, metrics_world_model, metrics_actor_critic = {}, {}, {}\n",
|
||||
"\n",
|
||||
"cfg_tokenizer = self.cfg.training.tokenizer\n",
|
||||
"cfg_world_model = self.cfg.training.world_model\n",
|
||||
"cfg_actor_critic = self.cfg.training.actor_critic\n",
|
||||
"\n",
|
||||
"# if epoch > cfg_tokenizer.start_after_epochs:\n",
|
||||
"# metrics_tokenizer = self.train_component(self.agent.tokenizer, self.optimizer_tokenizer, sequence_length=1, sample_from_start=True, **cfg_tokenizer)\n",
|
||||
"# self.agent.tokenizer.eval()\n",
|
||||
"\n",
|
||||
"# if epoch > cfg_world_model.start_after_epochs:\n",
|
||||
"# 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)\n",
|
||||
"# self.agent.world_model.eval()\n",
|
||||
"\n",
|
||||
"# if epoch > cfg_actor_critic.start_after_epochs:\n",
|
||||
"# 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)\n",
|
||||
"# self.agent.actor_critic.eval()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"from torchinfo import summary\n",
|
||||
"import torch\n",
|
||||
"from einops import rearrange\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Directly benchmark models"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"tokenizer = self.agent.tokenizer\n",
|
||||
"world_model = self.agent.world_model\n",
|
||||
"actor_critic = self.agent.actor_critic\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"4"
|
||||
]
|
||||
},
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"batch_num_samples = cfg.training.world_model.batch_num_samples\n",
|
||||
"sequence_length = cfg.common.sequence_length\n",
|
||||
"sample_from_start = False\n",
|
||||
"# train_dataset = instantiate(cfg.datasets.train)\n",
|
||||
"batch_num_samples\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"batch = self.train_dataset.sample_batch(batch_num_samples, sequence_length, sample_from_start)\n",
|
||||
"batch = {k: v.to(self.device) for k, v in batch.items()}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"CPU times: user 190 ms, sys: 4.87 ms, total: 194 ms\n",
|
||||
"Wall time: 195 ms\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<src.utils.LossWithIntermediateLosses at 0x7f97ed202a30>"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%time\n",
|
||||
"self.agent.world_model.compute_loss(batch, tokenizer=self.agent.tokenizer)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"CPU times: user 9.35 s, sys: 17.3 ms, total: 9.37 s\n",
|
||||
"Wall time: 9.37 s\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<src.utils.LossWithIntermediateLosses at 0x7f97f274beb0>"
|
||||
]
|
||||
},
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%time\n",
|
||||
"# TODO: why is this so slow?\n",
|
||||
"cfg_actor_critic = self.cfg.training.actor_critic\n",
|
||||
"self.agent.actor_critic.compute_loss(batch, tokenizer=self.agent.tokenizer, world_model=self.agent.world_model, **cfg_actor_critic)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"CPU times: user 9.11 s, sys: 15.5 ms, total: 9.13 s\n",
|
||||
"Wall time: 9.13 s\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%time\n",
|
||||
"# is this the slow part... yes. damn\n",
|
||||
"actor_critic.imagine(batch, tokenizer, world_model, horizon=10);\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # takes 0.1 s, fast\n",
|
||||
"# wm_env = WorldModelEnv(tokenizer, world_model, device)\n",
|
||||
"# wm_env\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 13,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"torch.Size([4, 3, 64, 64])\n",
|
||||
"CPU times: user 105 ms, sys: 207 µs, total: 105 ms\n",
|
||||
"Wall time: 105 ms\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%time\n",
|
||||
"# this takes 0.1 seconds and is run 10+ time. So 1 second. Hmm\n",
|
||||
"from src.envs.world_model_env import WorldModelEnv, Categorical\n",
|
||||
"initial_observations = batch['observations']\n",
|
||||
"\n",
|
||||
"# get the right obs\n",
|
||||
"wm_env = WorldModelEnv(self.agent.tokenizer, self.agent.world_model, self.device)\n",
|
||||
"obs = wm_env.reset_from_initial_observations(initial_observations[:, -1])\n",
|
||||
"print(obs.shape)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# make sure hidden states are right\n",
|
||||
"self.agent.actor_critic.reset(obs.shape[0])\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 14,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import gc\n",
|
||||
"gc.collect()\n",
|
||||
"torch.cuda.empty_cache()\n",
|
||||
"# obs\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 15,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"CPU times: user 1.6 ms, sys: 309 µs, total: 1.9 ms\n",
|
||||
"Wall time: 1.72 ms\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%time\n",
|
||||
"# 700us\n",
|
||||
"# fast, executed 10+ times\n",
|
||||
"outputs_ac = actor_critic(obs)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"torch.Size([4, 1, 17])"
|
||||
]
|
||||
},
|
||||
"execution_count": 16,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"outputs_ac.logits_actions.shape\n",
|
||||
"# action_token.shape\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 17,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# %%timeit\n",
|
||||
"# slow! takes 1s, executed 10+ times this is the culprit, not the lstm. hmm\n",
|
||||
"k=3\n",
|
||||
"horizon = 6\n",
|
||||
"action_token = Categorical(logits=outputs_ac.logits_actions).sample()\n",
|
||||
"obs, reward, done, _ = wm_env.step(action_token, should_predict_next_obs=(k < horizon - 1))\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"66.5 ms ± 1.53 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%timeit\n",
|
||||
"# 62ms\n",
|
||||
"# this is the slow part again. no grad and eval don't hepl\n",
|
||||
"outputs_wm = world_model(action_token, past_keys_values=wm_env.keys_values_wm)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"num_steps=1\n",
|
||||
"prev_steps=0\n",
|
||||
"sequences = world_model.embedder(action_token, num_steps, prev_steps) + world_model.pos_emb(prev_steps + torch.arange(num_steps, device=action_token.device))\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"ename": "AssertionError",
|
||||
"evalue": "",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||||
"\u001b[0;31mAssertionError\u001b[0m Traceback (most recent call last)",
|
||||
"\u001b[1;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/notebooks/01_debug_models.ipynb Cell 24\u001b[0m line \u001b[0;36m1\n\u001b[0;32m----> <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/notebooks/01_debug_models.ipynb#X64sZmlsZQ%3D%3D?line=0'>1</a>\u001b[0m get_ipython()\u001b[39m.\u001b[39;49mrun_cell_magic(\u001b[39m'\u001b[39;49m\u001b[39mtimeit\u001b[39;49m\u001b[39m'\u001b[39;49m, \u001b[39m'\u001b[39;49m\u001b[39m'\u001b[39;49m, \u001b[39m\"\u001b[39;49m\u001b[39m# ofc it\u001b[39;49m\u001b[39m'\u001b[39;49m\u001b[39ms the transformer that\u001b[39;49m\u001b[39m'\u001b[39;49m\u001b[39ms slow. I guess we just call it was more than during training\u001b[39;49m\u001b[39m\\n\u001b[39;49;00m\u001b[39mpast_keys_values = wm_env.keys_values_wm\u001b[39;49m\u001b[39m\\n\u001b[39;49;00m\u001b[39mx = world_model.transformer(sequences, past_keys_values)\u001b[39;49m\u001b[39m\\n\u001b[39;49;00m\u001b[39m\"\u001b[39;49m)\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/interactiveshell.py:2515\u001b[0m, in \u001b[0;36mInteractiveShell.run_cell_magic\u001b[0;34m(self, magic_name, line, cell)\u001b[0m\n\u001b[1;32m 2513\u001b[0m \u001b[39mwith\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39mbuiltin_trap:\n\u001b[1;32m 2514\u001b[0m args \u001b[39m=\u001b[39m (magic_arg_s, cell)\n\u001b[0;32m-> 2515\u001b[0m result \u001b[39m=\u001b[39m fn(\u001b[39m*\u001b[39;49margs, \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mkwargs)\n\u001b[1;32m 2517\u001b[0m \u001b[39m# The code below prevents the output from being displayed\u001b[39;00m\n\u001b[1;32m 2518\u001b[0m \u001b[39m# when using magics with decorator @output_can_be_silenced\u001b[39;00m\n\u001b[1;32m 2519\u001b[0m \u001b[39m# when the last Python token in the expression is a ';'.\u001b[39;00m\n\u001b[1;32m 2520\u001b[0m \u001b[39mif\u001b[39;00m \u001b[39mgetattr\u001b[39m(fn, magic\u001b[39m.\u001b[39mMAGIC_OUTPUT_CAN_BE_SILENCED, \u001b[39mFalse\u001b[39;00m):\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/magics/execution.py:1189\u001b[0m, in \u001b[0;36mExecutionMagics.timeit\u001b[0;34m(self, line, cell, local_ns)\u001b[0m\n\u001b[1;32m 1186\u001b[0m \u001b[39mif\u001b[39;00m time_number \u001b[39m>\u001b[39m\u001b[39m=\u001b[39m \u001b[39m0.2\u001b[39m:\n\u001b[1;32m 1187\u001b[0m \u001b[39mbreak\u001b[39;00m\n\u001b[0;32m-> 1189\u001b[0m all_runs \u001b[39m=\u001b[39m timer\u001b[39m.\u001b[39;49mrepeat(repeat, number)\n\u001b[1;32m 1190\u001b[0m best \u001b[39m=\u001b[39m \u001b[39mmin\u001b[39m(all_runs) \u001b[39m/\u001b[39m number\n\u001b[1;32m 1191\u001b[0m worst \u001b[39m=\u001b[39m \u001b[39mmax\u001b[39m(all_runs) \u001b[39m/\u001b[39m number\n",
|
||||
"File \u001b[0;32m~/miniforge3/lib/python3.9/timeit.py:205\u001b[0m, in \u001b[0;36mTimer.repeat\u001b[0;34m(self, repeat, number)\u001b[0m\n\u001b[1;32m 203\u001b[0m r \u001b[39m=\u001b[39m []\n\u001b[1;32m 204\u001b[0m \u001b[39mfor\u001b[39;00m i \u001b[39min\u001b[39;00m \u001b[39mrange\u001b[39m(repeat):\n\u001b[0;32m--> 205\u001b[0m t \u001b[39m=\u001b[39m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49mtimeit(number)\n\u001b[1;32m 206\u001b[0m r\u001b[39m.\u001b[39mappend(t)\n\u001b[1;32m 207\u001b[0m \u001b[39mreturn\u001b[39;00m r\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/magics/execution.py:173\u001b[0m, in \u001b[0;36mTimer.timeit\u001b[0;34m(self, number)\u001b[0m\n\u001b[1;32m 171\u001b[0m gc\u001b[39m.\u001b[39mdisable()\n\u001b[1;32m 172\u001b[0m \u001b[39mtry\u001b[39;00m:\n\u001b[0;32m--> 173\u001b[0m timing \u001b[39m=\u001b[39m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49minner(it, \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49mtimer)\n\u001b[1;32m 174\u001b[0m \u001b[39mfinally\u001b[39;00m:\n\u001b[1;32m 175\u001b[0m \u001b[39mif\u001b[39;00m gcold:\n",
|
||||
"File \u001b[0;32m<magic-timeit>:3\u001b[0m, in \u001b[0;36minner\u001b[0;34m(_it, _timer)\u001b[0m\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/nn/modules/module.py:1518\u001b[0m, in \u001b[0;36mModule._wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m 1516\u001b[0m \u001b[39mreturn\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_compiled_call_impl(\u001b[39m*\u001b[39margs, \u001b[39m*\u001b[39m\u001b[39m*\u001b[39mkwargs) \u001b[39m# type: ignore[misc]\u001b[39;00m\n\u001b[1;32m 1517\u001b[0m \u001b[39melse\u001b[39;00m:\n\u001b[0;32m-> 1518\u001b[0m \u001b[39mreturn\u001b[39;00m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49m_call_impl(\u001b[39m*\u001b[39;49margs, \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mkwargs)\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/nn/modules/module.py:1527\u001b[0m, in \u001b[0;36mModule._call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m 1522\u001b[0m \u001b[39m# If we don't have any hooks, we want to skip the rest of the logic in\u001b[39;00m\n\u001b[1;32m 1523\u001b[0m \u001b[39m# this function, and just call forward.\u001b[39;00m\n\u001b[1;32m 1524\u001b[0m \u001b[39mif\u001b[39;00m \u001b[39mnot\u001b[39;00m (\u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_backward_hooks \u001b[39mor\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_backward_pre_hooks \u001b[39mor\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_forward_hooks \u001b[39mor\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_forward_pre_hooks\n\u001b[1;32m 1525\u001b[0m \u001b[39mor\u001b[39;00m _global_backward_pre_hooks \u001b[39mor\u001b[39;00m _global_backward_hooks\n\u001b[1;32m 1526\u001b[0m \u001b[39mor\u001b[39;00m _global_forward_hooks \u001b[39mor\u001b[39;00m _global_forward_pre_hooks):\n\u001b[0;32m-> 1527\u001b[0m \u001b[39mreturn\u001b[39;00m forward_call(\u001b[39m*\u001b[39;49margs, \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mkwargs)\n\u001b[1;32m 1529\u001b[0m \u001b[39mtry\u001b[39;00m:\n\u001b[1;32m 1530\u001b[0m result \u001b[39m=\u001b[39m \u001b[39mNone\u001b[39;00m\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/transformer.py:69\u001b[0m, in \u001b[0;36mTransformer.forward\u001b[0;34m(self, sequences, past_keys_values)\u001b[0m\n\u001b[1;32m 66\u001b[0m \u001b[39m# k_size = (x.shape[0], x.shape[1], x.shape[1], 1)\u001b[39;00m\n\u001b[1;32m 67\u001b[0m \u001b[39m# v_size = past_keys_values[0]._v_cache._cache.size()\u001b[39;00m\n\u001b[1;32m 68\u001b[0m v_size \u001b[39m=\u001b[39m (k_size[\u001b[39m0\u001b[39m], k_size[\u001b[39m1\u001b[39m], x\u001b[39m.\u001b[39mshape[\u001b[39m1\u001b[39m], k_size[\u001b[39m3\u001b[39m])\n\u001b[0;32m---> 69\u001b[0m past_keys_values[\u001b[39m0\u001b[39;49m]\u001b[39m.\u001b[39;49mupdate(torch\u001b[39m.\u001b[39;49mrand(v_size), torch\u001b[39m.\u001b[39;49mrand(v_size))\n\u001b[1;32m 70\u001b[0m \u001b[39mreturn\u001b[39;00m x\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/kv_caching.py:59\u001b[0m, in \u001b[0;36mKVCache.update\u001b[0;34m(self, k, v)\u001b[0m\n\u001b[1;32m 58\u001b[0m \u001b[39mdef\u001b[39;00m \u001b[39mupdate\u001b[39m(\u001b[39mself\u001b[39m, k: torch\u001b[39m.\u001b[39mTensor, v: torch\u001b[39m.\u001b[39mTensor):\n\u001b[0;32m---> 59\u001b[0m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49m_k_cache\u001b[39m.\u001b[39;49mupdate(k)\n\u001b[1;32m 60\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_v_cache\u001b[39m.\u001b[39mupdate(v)\n",
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/kv_caching.py:33\u001b[0m, in \u001b[0;36mCache.update\u001b[0;34m(self, x)\u001b[0m\n\u001b[1;32m 31\u001b[0m \u001b[39mdef\u001b[39;00m \u001b[39mupdate\u001b[39m(\u001b[39mself\u001b[39m, x: torch\u001b[39m.\u001b[39mTensor) \u001b[39m-\u001b[39m\u001b[39m>\u001b[39m \u001b[39mNone\u001b[39;00m:\n\u001b[1;32m 32\u001b[0m \u001b[39massert\u001b[39;00m (x\u001b[39m.\u001b[39mndim \u001b[39m==\u001b[39m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_cache\u001b[39m.\u001b[39mndim) \u001b[39mand\u001b[39;00m \u001b[39mall\u001b[39m([x\u001b[39m.\u001b[39msize(i) \u001b[39m==\u001b[39m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_cache\u001b[39m.\u001b[39msize(i) \u001b[39mfor\u001b[39;00m i \u001b[39min\u001b[39;00m (\u001b[39m0\u001b[39m, \u001b[39m1\u001b[39m, \u001b[39m3\u001b[39m)])\n\u001b[0;32m---> 33\u001b[0m \u001b[39massert\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_size \u001b[39m+\u001b[39m x\u001b[39m.\u001b[39msize(\u001b[39m2\u001b[39m) \u001b[39m<\u001b[39m\u001b[39m=\u001b[39m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_cache\u001b[39m.\u001b[39mshape[\u001b[39m2\u001b[39m]\n\u001b[1;32m 34\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_cache \u001b[39m=\u001b[39m AssignWithoutInplaceCheck\u001b[39m.\u001b[39mapply(\u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_cache, x, \u001b[39m2\u001b[39m, \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_size, \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_size \u001b[39m+\u001b[39m x\u001b[39m.\u001b[39msize(\u001b[39m2\u001b[39m))\n\u001b[1;32m 35\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_size \u001b[39m+\u001b[39m\u001b[39m=\u001b[39m x\u001b[39m.\u001b[39msize(\u001b[39m2\u001b[39m)\n",
|
||||
"\u001b[0;31mAssertionError\u001b[0m: "
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"%%timeit\n",
|
||||
"# ofc it's the transformer that's slow. I guess we just call it was more than during training\n",
|
||||
"past_keys_values = wm_env.keys_values_wm\n",
|
||||
"x = world_model.transformer(sequences, past_keys_values)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# past_keys_values = wm_env.keys_values_wm\n",
|
||||
"# x = world_model.transformer(sequences, past_keys_values)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"%%timeit\n",
|
||||
"# ofc it's the transformer that's slow. I guess we just call it was more than during training\n",
|
||||
"past_keys_values = wm_env.keys_values_wm\n",
|
||||
"x = world_model.transformer(sequences, past_keys_values)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"logits_observations = world_model.head_observations(x, num_steps=num_steps, prev_steps=prev_steps)\n",
|
||||
"logits_rewards = world_model.head_rewards(x, num_steps=num_steps, prev_steps=prev_steps)\n",
|
||||
"logits_ends = world_model.head_ends(x, num_steps=num_steps, prev_steps=prev_steps)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Torchinfo model sizes\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"observations = self.agent.tokenizer.preprocess_input(rearrange(batch['observations'], 'b t c h w -> (b t) c h w'))\n",
|
||||
"# z, z_quantized, reconstructions = self.agent.tokenizer(observations, should_preprocess=False, should_postprocess=False)\n",
|
||||
"summary(self.agent.tokenizer, input_data=observations)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"\n",
|
||||
"with torch.no_grad():\n",
|
||||
" obs_tokens = self.agent.tokenizer.encode(batch['observations'], should_preprocess=True).tokens # (BL, K)\n",
|
||||
"\n",
|
||||
"act_tokens = rearrange(batch['actions'], 'b l -> b l 1')\n",
|
||||
"tokens = rearrange(torch.cat((obs_tokens, act_tokens), dim=2), 'b l k1 -> b (l k1)') # \n",
|
||||
"\n",
|
||||
"summary(self.agent.world_model, input_data=tokens)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from src.envs.world_model_env import WorldModelEnv\n",
|
||||
"initial_observations = batch['observations']\n",
|
||||
"\n",
|
||||
"# get the right obs\n",
|
||||
"wm_env = WorldModelEnv(self.agent.tokenizer, self.agent.world_model, self.device)\n",
|
||||
"obs = wm_env.reset_from_initial_observations(initial_observations[:, -1])\n",
|
||||
"obs.shape\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# make sure hidden states are right\n",
|
||||
"self.agent.actor_critic.reset(obs.shape[0])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"\n",
|
||||
"from torchinfo import summary\n",
|
||||
"summary(self.agent.actor_critic, input_data=obs)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Debug env\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import minihack\n",
|
||||
"env = gym.make(\"MiniHack-River-v0\", observation_keys=(\"pixel_crop\", \"pixel\", 'blstats', 'message'))\n",
|
||||
"env.reset() # each reset generates a new environment instance\n",
|
||||
"obs, reward, end, info = env.step(1) # move agent '@' north\n",
|
||||
"print(obs['pixel_crop'].shape)\n",
|
||||
"plt.imshow(obs['pixel_crop'])\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"print(obs['pixel'].shape)\n",
|
||||
"plt.imshow(obs['pixel'])\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # plt.imshow(obs['glyphs_crop'])\n",
|
||||
"# obs['glyphs_crop'].shape\n",
|
||||
"# obs['blstats']\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import minihack\n",
|
||||
"env = gym.make(\"MiniHack-Room-5x5-v0\", observation_keys=(\"pixel_crop\", \"pixel\", 'blstats', 'message'))\n",
|
||||
"env.reset() # each reset generates a new environment instance\n",
|
||||
"obs, reward, end, info = env.step(1) # move agent '@' north\n",
|
||||
"print(obs['pixel_crop'].shape)\n",
|
||||
"plt.imshow(obs['pixel_crop'])\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"print(obs['pixel'].shape)\n",
|
||||
"plt.imshow(obs['pixel'])\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import minihack\n",
|
||||
"import crafter\n",
|
||||
"env = gym.make(\"CrafterReward-v1\")\n",
|
||||
"env.reset() # each reset generates a new environment instance\n",
|
||||
"obs, reward, end, info = env.step(1) # move agent '@' north\n",
|
||||
"print(obs.shape)\n",
|
||||
"plt.imshow(obs)\n",
|
||||
"plt.show()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": ".venv",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.9.16"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
Generated
+199
-1
@@ -441,6 +441,26 @@ mypy = ["contourpy[bokeh,docs]", "docutils-stubs", "mypy (==1.6.1)", "types-Pill
|
||||
test = ["Pillow", "contourpy[test-no-images]", "matplotlib"]
|
||||
test-no-images = ["pytest", "pytest-cov", "pytest-xdist", "wurlitzer"]
|
||||
|
||||
[[package]]
|
||||
name = "crafter"
|
||||
version = "1.8.2"
|
||||
description = "Open world survival game for reinforcement learning."
|
||||
optional = false
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "crafter-1.8.2.tar.gz", hash = "sha256:4a142c291aa0b137c0808890381b38803876a5811b362801b8cc669faa350def"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
imageio = "*"
|
||||
numpy = "*"
|
||||
opensimplex = "*"
|
||||
pillow = "*"
|
||||
"ruamel.yaml" = "*"
|
||||
|
||||
[package.extras]
|
||||
gui = ["pygame"]
|
||||
|
||||
[[package]]
|
||||
name = "cycler"
|
||||
version = "0.12.1"
|
||||
@@ -794,6 +814,37 @@ files = [
|
||||
{file = "idna-3.4.tar.gz", hash = "sha256:814f528e8dead7d329833b91c5faa87d60bf71824cd12a7530b5526063d02cb4"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "imageio"
|
||||
version = "2.31.5"
|
||||
description = "Library for reading and writing a wide range of image, video, scientific, and volumetric data formats."
|
||||
optional = false
|
||||
python-versions = ">=3.8"
|
||||
files = [
|
||||
{file = "imageio-2.31.5-py3-none-any.whl", hash = "sha256:97f68e12ba676f2f4b541684ed81f7f3370dc347e8321bc68ee34d37b2dbac9f"},
|
||||
{file = "imageio-2.31.5.tar.gz", hash = "sha256:d8e53f9cd4054880276a3dac0a28c85ba7874084856a55a0294a8ae6ed7f3a8e"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
numpy = "*"
|
||||
pillow = ">=8.3.2"
|
||||
|
||||
[package.extras]
|
||||
all-plugins = ["astropy", "av", "imageio-ffmpeg", "psutil", "tifffile"]
|
||||
all-plugins-pypy = ["av", "imageio-ffmpeg", "psutil", "tifffile"]
|
||||
build = ["wheel"]
|
||||
dev = ["black", "flake8", "fsspec[github]", "pytest", "pytest-cov"]
|
||||
docs = ["numpydoc", "pydata-sphinx-theme", "sphinx (<6)"]
|
||||
ffmpeg = ["imageio-ffmpeg", "psutil"]
|
||||
fits = ["astropy"]
|
||||
full = ["astropy", "av", "black", "flake8", "fsspec[github]", "gdal", "imageio-ffmpeg", "itk", "numpydoc", "psutil", "pydata-sphinx-theme", "pytest", "pytest-cov", "sphinx (<6)", "tifffile", "wheel"]
|
||||
gdal = ["gdal"]
|
||||
itk = ["itk"]
|
||||
linting = ["black", "flake8"]
|
||||
pyav = ["av"]
|
||||
test = ["fsspec[github]", "pytest", "pytest-cov"]
|
||||
tifffile = ["tifffile"]
|
||||
|
||||
[[package]]
|
||||
name = "importlib-metadata"
|
||||
version = "6.8.0"
|
||||
@@ -1244,6 +1295,28 @@ files = [
|
||||
[package.dependencies]
|
||||
traitlets = "*"
|
||||
|
||||
[[package]]
|
||||
name = "minihack"
|
||||
version = "0.1.5"
|
||||
description = "MiniHack The Planet: A Sandbox for Open-Ended Reinforcement Learning Research"
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
files = [
|
||||
{file = "minihack-0.1.5.tar.gz", hash = "sha256:b732a05eed9b70cc48735d63fd0c34498611be1342f4bd45a3ad128367d8c97a"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
gym = ">=0.15,<=0.23"
|
||||
nle = "0.9.0"
|
||||
numpy = ">=1.16"
|
||||
|
||||
[package.extras]
|
||||
all = ["black (>=19.10b0)", "flake8 (>=3.7)", "flake8-bugbear (>=20.1)", "hydra-colorlog (>=1.0.0)", "hydra-colorlog (>=1.0.0)", "hydra-core (>=1.0.0)", "hydra-core (>=1.0.0)", "hydra-submitit-launcher (>=1.1.1)", "hydra-submitit-launcher (>=1.1.1)", "inflect", "myst-parser (==0.15.1)", "nbsphinx (==0.8.6)", "pre-commit (>=2.0.1)", "pytest (>=5.3)", "pytest-benchmark (>=3.1.0)", "pyyaml", "ray[default] (==1.3.0)", "ray[rllib] (==1.3.0)", "sphinx (==4.0.2)", "sphinx-rtd-theme (==1.0.0)", "stanza", "torch (>=1.3.1)", "torch (>=1.3.1)", "wandb (>=0.10.31)", "wandb (>=0.10.31)"]
|
||||
dev = ["black (>=19.10b0)", "flake8 (>=3.7)", "flake8-bugbear (>=20.1)", "myst-parser (==0.15.1)", "nbsphinx (==0.8.6)", "pre-commit (>=2.0.1)", "pytest (>=5.3)", "pytest-benchmark (>=3.1.0)", "sphinx (==4.0.2)", "sphinx-rtd-theme (==1.0.0)"]
|
||||
polybeast = ["hydra-colorlog (>=1.0.0)", "hydra-core (>=1.0.0)", "hydra-submitit-launcher (>=1.1.1)", "pyyaml", "torch (>=1.3.1)", "wandb (>=0.10.31)"]
|
||||
rllib = ["hydra-colorlog (>=1.0.0)", "hydra-core (>=1.0.0)", "hydra-submitit-launcher (>=1.1.1)", "ray[default] (==1.3.0)", "ray[rllib] (==1.3.0)", "torch (>=1.3.1)", "wandb (>=0.10.31)"]
|
||||
wiki = ["inflect", "stanza"]
|
||||
|
||||
[[package]]
|
||||
name = "mpmath"
|
||||
version = "1.3.0"
|
||||
@@ -1290,6 +1363,26 @@ doc = ["nb2plots (>=0.7)", "nbconvert (<7.9)", "numpydoc (>=1.6)", "pillow (>=9.
|
||||
extra = ["lxml (>=4.6)", "pydot (>=1.4.2)", "pygraphviz (>=1.11)", "sympy (>=1.10)"]
|
||||
test = ["pytest (>=7.2)", "pytest-cov (>=4.0)"]
|
||||
|
||||
[[package]]
|
||||
name = "nle"
|
||||
version = "0.9.0"
|
||||
description = "The NetHack Learning Environment (NLE): a reinforcement learning environment based on NetHack"
|
||||
optional = false
|
||||
python-versions = ">=3.5"
|
||||
files = [
|
||||
{file = "nle-0.9.0.tar.gz", hash = "sha256:a98644bdd547017cdde9fdf589b245da98ff8753327285e949d506d9006f10d4"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
gym = ">=0.15,<=0.23"
|
||||
numpy = ">=1.16"
|
||||
pybind11 = ">=2.2"
|
||||
|
||||
[package.extras]
|
||||
agent = ["torch (>=1.3.1)"]
|
||||
all = ["black (>=19.10b0)", "cmake_format (>=0.6.10)", "flake8 (>=3.7)", "flake8-bugbear (>=20.1)", "memory-profiler (>=0.60.0)", "pre-commit (>=2.0.1)", "pytest (>=6.2.5)", "pytest-benchmark (>=3.4.1)", "sphinx (>=2.4.4)", "sphinx-rtd-theme (==0.4.3)", "torch (>=1.3.1)"]
|
||||
dev = ["black (>=19.10b0)", "cmake_format (>=0.6.10)", "flake8 (>=3.7)", "flake8-bugbear (>=20.1)", "memory-profiler (>=0.60.0)", "pre-commit (>=2.0.1)", "pytest (>=6.2.5)", "pytest-benchmark (>=3.4.1)", "sphinx (>=2.4.4)", "sphinx-rtd-theme (==0.4.3)"]
|
||||
|
||||
[[package]]
|
||||
name = "numpy"
|
||||
version = "1.26.1"
|
||||
@@ -1371,6 +1464,20 @@ numpy = [
|
||||
{version = ">=1.19.3", markers = "platform_system == \"Linux\" and platform_machine == \"aarch64\" and python_version >= \"3.8\" and python_version < \"3.10\" or python_version > \"3.9\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_system != \"Darwin\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_machine != \"arm64\" and python_version < \"3.10\""},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "opensimplex"
|
||||
version = "0.4.5"
|
||||
description = "OpenSimplex is a noise generation function like Perlin or Simplex noise, but better."
|
||||
optional = false
|
||||
python-versions = ">=3.8"
|
||||
files = [
|
||||
{file = "opensimplex-0.4.5-py3-none-any.whl", hash = "sha256:5e34f2b6f7e2d3e798d7f060b45cb1f92c02b68b6fc2626ac2a36bd53ec6c773"},
|
||||
{file = "opensimplex-0.4.5.tar.gz", hash = "sha256:c390cf70dea97b32bd1a49ba6781e84f48dc93b1cc6c06f36d9d44c548299f90"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
numpy = ">=1.22"
|
||||
|
||||
[[package]]
|
||||
name = "packaging"
|
||||
version = "23.2"
|
||||
@@ -1645,6 +1752,20 @@ files = [
|
||||
[package.extras]
|
||||
tests = ["pytest"]
|
||||
|
||||
[[package]]
|
||||
name = "pybind11"
|
||||
version = "2.11.1"
|
||||
description = "Seamless operability between C++11 and Python"
|
||||
optional = false
|
||||
python-versions = ">=3.6"
|
||||
files = [
|
||||
{file = "pybind11-2.11.1-py3-none-any.whl", hash = "sha256:33cdd02a6453380dd71cc70357ce388ad1ee8d32bd0e38fc22b273d050aa29b3"},
|
||||
{file = "pybind11-2.11.1.tar.gz", hash = "sha256:00cd59116a6e8155aecd9174f37ba299d1d397ed4a6b86ac1dfe01b3e40f2cc4"},
|
||||
]
|
||||
|
||||
[package.extras]
|
||||
global = ["pybind11-global (==2.11.1)"]
|
||||
|
||||
[[package]]
|
||||
name = "pycparser"
|
||||
version = "2.21"
|
||||
@@ -2069,6 +2190,83 @@ urllib3 = ">=1.21.1,<3"
|
||||
socks = ["PySocks (>=1.5.6,!=1.5.7)"]
|
||||
use-chardet-on-py3 = ["chardet (>=3.0.2,<6)"]
|
||||
|
||||
[[package]]
|
||||
name = "ruamel-yaml"
|
||||
version = "0.18.5"
|
||||
description = "ruamel.yaml is a YAML parser/emitter that supports roundtrip preservation of comments, seq/map flow style, and map key order"
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
files = [
|
||||
{file = "ruamel.yaml-0.18.5-py3-none-any.whl", hash = "sha256:a013ac02f99a69cdd6277d9664689eb1acba07069f912823177c5eced21a6ada"},
|
||||
{file = "ruamel.yaml-0.18.5.tar.gz", hash = "sha256:61917e3a35a569c1133a8f772e1226961bf5a1198bea7e23f06a0841dea1ab0e"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
"ruamel.yaml.clib" = {version = ">=0.2.7", markers = "platform_python_implementation == \"CPython\" and python_version < \"3.13\""}
|
||||
|
||||
[package.extras]
|
||||
docs = ["mercurial (>5.7)", "ryd"]
|
||||
jinja2 = ["ruamel.yaml.jinja2 (>=0.2)"]
|
||||
|
||||
[[package]]
|
||||
name = "ruamel-yaml-clib"
|
||||
version = "0.2.8"
|
||||
description = "C version of reader, parser and emitter for ruamel.yaml derived from libyaml"
|
||||
optional = false
|
||||
python-versions = ">=3.6"
|
||||
files = [
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:b42169467c42b692c19cf539c38d4602069d8c1505e97b86387fcf7afb766e1d"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-macosx_13_0_arm64.whl", hash = "sha256:07238db9cbdf8fc1e9de2489a4f68474e70dffcb32232db7c08fa61ca0c7c462"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.manylinux_2_24_x86_64.whl", hash = "sha256:fff3573c2db359f091e1589c3d7c5fc2f86f5bdb6f24252c2d8e539d4e45f412"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-manylinux_2_24_aarch64.whl", hash = "sha256:aa2267c6a303eb483de8d02db2871afb5c5fc15618d894300b88958f729ad74f"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-musllinux_1_1_i686.whl", hash = "sha256:840f0c7f194986a63d2c2465ca63af8ccbbc90ab1c6001b1978f05119b5e7334"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-musllinux_1_1_x86_64.whl", hash = "sha256:024cfe1fc7c7f4e1aff4a81e718109e13409767e4f871443cbff3dba3578203d"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-win32.whl", hash = "sha256:c69212f63169ec1cfc9bb44723bf2917cbbd8f6191a00ef3410f5a7fe300722d"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp310-cp310-win_amd64.whl", hash = "sha256:cabddb8d8ead485e255fe80429f833172b4cadf99274db39abc080e068cbcc31"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-macosx_10_9_universal2.whl", hash = "sha256:bef08cd86169d9eafb3ccb0a39edb11d8e25f3dae2b28f5c52fd997521133069"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-macosx_13_0_arm64.whl", hash = "sha256:b16420e621d26fdfa949a8b4b47ade8810c56002f5389970db4ddda51dbff248"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.manylinux_2_24_x86_64.whl", hash = "sha256:25c515e350e5b739842fc3228d662413ef28f295791af5e5110b543cf0b57d9b"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-manylinux_2_24_aarch64.whl", hash = "sha256:1707814f0d9791df063f8c19bb51b0d1278b8e9a2353abbb676c2f685dee6afe"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-musllinux_1_1_i686.whl", hash = "sha256:46d378daaac94f454b3a0e3d8d78cafd78a026b1d71443f4966c696b48a6d899"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:09b055c05697b38ecacb7ac50bdab2240bfca1a0c4872b0fd309bb07dc9aa3a9"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-win32.whl", hash = "sha256:53a300ed9cea38cf5a2a9b069058137c2ca1ce658a874b79baceb8f892f915a7"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp311-cp311-win_amd64.whl", hash = "sha256:c2a72e9109ea74e511e29032f3b670835f8a59bbdc9ce692c5b4ed91ccf1eedb"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-macosx_10_9_universal2.whl", hash = "sha256:ebc06178e8821efc9692ea7544aa5644217358490145629914d8020042c24aa1"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-macosx_13_0_arm64.whl", hash = "sha256:edaef1c1200c4b4cb914583150dcaa3bc30e592e907c01117c08b13a07255ec2"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:d176b57452ab5b7028ac47e7b3cf644bcfdc8cacfecf7e71759f7f51a59e5c92"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-manylinux_2_24_aarch64.whl", hash = "sha256:1dc67314e7e1086c9fdf2680b7b6c2be1c0d8e3a8279f2e993ca2a7545fecf62"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-musllinux_1_1_i686.whl", hash = "sha256:3213ece08ea033eb159ac52ae052a4899b56ecc124bb80020d9bbceeb50258e9"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:aab7fd643f71d7946f2ee58cc88c9b7bfc97debd71dcc93e03e2d174628e7e2d"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-win32.whl", hash = "sha256:5c365d91c88390c8d0a8545df0b5857172824b1c604e867161e6b3d59a827eaa"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp312-cp312-win_amd64.whl", hash = "sha256:1758ce7d8e1a29d23de54a16ae867abd370f01b5a69e1a3ba75223eaa3ca1a1b"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp36-cp36m-manylinux_2_5_x86_64.manylinux1_x86_64.whl", hash = "sha256:a5aa27bad2bb83670b71683aae140a1f52b0857a2deff56ad3f6c13a017a26ed"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-macosx_10_9_x86_64.whl", hash = "sha256:c58ecd827313af6864893e7af0a3bb85fd529f862b6adbefe14643947cfe2942"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-macosx_12_0_arm64.whl", hash = "sha256:f481f16baec5290e45aebdc2a5168ebc6d35189ae6fea7a58787613a25f6e875"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-manylinux_2_24_aarch64.whl", hash = "sha256:77159f5d5b5c14f7c34073862a6b7d34944075d9f93e681638f6d753606c6ce6"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-manylinux_2_5_x86_64.manylinux1_x86_64.whl", hash = "sha256:7f67a1ee819dc4562d444bbafb135832b0b909f81cc90f7aa00260968c9ca1b3"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-musllinux_1_1_i686.whl", hash = "sha256:4ecbf9c3e19f9562c7fdd462e8d18dd902a47ca046a2e64dba80699f0b6c09b7"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-musllinux_1_1_x86_64.whl", hash = "sha256:87ea5ff66d8064301a154b3933ae406b0863402a799b16e4a1d24d9fbbcbe0d3"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-win32.whl", hash = "sha256:75e1ed13e1f9de23c5607fe6bd1aeaae21e523b32d83bb33918245361e9cc51b"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp37-cp37m-win_amd64.whl", hash = "sha256:3f215c5daf6a9d7bbed4a0a4f760f3113b10e82ff4c5c44bec20a68c8014f675"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-macosx_10_9_x86_64.whl", hash = "sha256:1b617618914cb00bf5c34d4357c37aa15183fa229b24767259657746c9077615"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-macosx_12_0_arm64.whl", hash = "sha256:a6a9ffd280b71ad062eae53ac1659ad86a17f59a0fdc7699fd9be40525153337"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-manylinux_2_24_aarch64.whl", hash = "sha256:305889baa4043a09e5b76f8e2a51d4ffba44259f6b4c72dec8ca56207d9c6fe1"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-manylinux_2_5_x86_64.manylinux1_x86_64.whl", hash = "sha256:700e4ebb569e59e16a976857c8798aee258dceac7c7d6b50cab63e080058df91"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-musllinux_1_1_i686.whl", hash = "sha256:e2b4c44b60eadec492926a7270abb100ef9f72798e18743939bdbf037aab8c28"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-musllinux_1_1_x86_64.whl", hash = "sha256:e79e5db08739731b0ce4850bed599235d601701d5694c36570a99a0c5ca41a9d"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-win32.whl", hash = "sha256:955eae71ac26c1ab35924203fda6220f84dce57d6d7884f189743e2abe3a9fbe"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp38-cp38-win_amd64.whl", hash = "sha256:56f4252222c067b4ce51ae12cbac231bce32aee1d33fbfc9d17e5b8d6966c312"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-macosx_10_9_x86_64.whl", hash = "sha256:03d1162b6d1df1caa3a4bd27aa51ce17c9afc2046c31b0ad60a0a96ec22f8001"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-macosx_12_0_arm64.whl", hash = "sha256:bba64af9fa9cebe325a62fa398760f5c7206b215201b0ec825005f1b18b9bccf"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-manylinux_2_24_aarch64.whl", hash = "sha256:a1a45e0bb052edf6a1d3a93baef85319733a888363938e1fc9924cb00c8df24c"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-manylinux_2_5_x86_64.manylinux1_x86_64.whl", hash = "sha256:da09ad1c359a728e112d60116f626cc9f29730ff3e0e7db72b9a2dbc2e4beed5"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-musllinux_1_1_i686.whl", hash = "sha256:184565012b60405d93838167f425713180b949e9d8dd0bbc7b49f074407c5a8b"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-musllinux_1_1_x86_64.whl", hash = "sha256:a75879bacf2c987c003368cf14bed0ffe99e8e85acfa6c0bfffc21a090f16880"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-win32.whl", hash = "sha256:84b554931e932c46f94ab306913ad7e11bba988104c5cff26d90d03f68258cd5"},
|
||||
{file = "ruamel.yaml.clib-0.2.8-cp39-cp39-win_amd64.whl", hash = "sha256:25ac8c08322002b06fa1d49d1646181f0b2c72f5cbc15a85e80b4c30a544bb15"},
|
||||
{file = "ruamel.yaml.clib-0.2.8.tar.gz", hash = "sha256:beb2e0404003de9a4cab9753a8805a8fe9320ee6673136ed7f04255fe60bb512"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ruff"
|
||||
version = "0.1.5"
|
||||
@@ -2940,4 +3138,4 @@ testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "p
|
||||
[metadata]
|
||||
lock-version = "2.0"
|
||||
python-versions = ">=3.9,<3.13"
|
||||
content-hash = "e2ad5cf8670b043bb7174a4c7f6e4c89644461eb1680e282b436541672e4e67c"
|
||||
content-hash = "d265b7789c918f4c2dc7d3db9ea80871958320fedc47576d0e0f78f327f42f6e"
|
||||
|
||||
@@ -28,6 +28,8 @@ torchvision = "^0.16.0"
|
||||
numpy = ">=1.18.0"
|
||||
gym = {version = "0.22.0", extras = ["accept-rom-license", "atari"]}
|
||||
scipy = "^1.11.3"
|
||||
crafter = "^1.8.2"
|
||||
minihack = "^0.1.5"
|
||||
|
||||
[[tool.poetry.source]]
|
||||
name = "pytorch"
|
||||
|
||||
@@ -232,3 +232,85 @@ Training tokenizer: 100%|██████████████████
|
||||
Training world_model: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [01:15<00:00, 2.64it/s]
|
||||
Training actor_critic: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20/20 [03:20<00:00, 10.05s/it]
|
||||
|
||||
|
||||
|
||||
- what about resume? oh we seem to have that although the code doesn't make sense https://hydra.cc/docs/tutorials/basic/running_your_app/working_directory/ https://hydra.cc/docs/1.2/upgrades/1.1_to_1.2/changes_to_job_working_dir/ see eval.py
|
||||
- [ ] but it's still too damn slow. what about bfloat16? using auto case?
|
||||
- why does it take so long? it would be nice to have a reproduction notebook
|
||||
- also the model might be to small now....
|
||||
|
||||
how to play
|
||||
```sh
|
||||
cd outputs/2023-11-17/07-59-44
|
||||
python scripts/play.sh
|
||||
```
|
||||
|
||||
|
||||
## Envs
|
||||
|
||||
tl:dr just use pong or breakout or crafter (1m steps)
|
||||
|
||||
for steps see [crafter paper](https://arxiv.org/pdf/2109.06780.pdf)
|
||||
|
||||
Nethack learning env. What's the obs size? 21x79 of glyphs (5991 possibilities) and 21 dim of stats
|
||||
- they use an lstm of 128. 5 layer conv
|
||||
- requires 1B steps
|
||||
-
|
||||
atari:
|
||||
- reqs 200M stpes
|
||||
|
||||
progcen:
|
||||
- 200M steps
|
||||
|
||||
minihack:
|
||||
- 2M steps for room 5xt
|
||||
- but needs editing to be atari compatible. e.g. 336 × 1264 × 3 pixels
|
||||
- pixel_crop 64,64,3 or 9x9 crop works!
|
||||
- lstm 256
|
||||
- The training on MiniHack’s Room-5x5 task for two million timesteps using our IMPALA baseline takes approximately 4:30 minutes (r
|
||||
|
||||
crafter
|
||||
- reqs 1M steps
|
||||
- "All agents trained for 1M environment steps in under 24 hours on a single GPU and we repeated the training for 10 random seeds per method. The training reward curves are included in Appendix "
|
||||
|
||||
|
||||
|
||||
|
||||
from https://arxiv.org/pdf/2111.09794.pdf
|
||||
There are several PCG state-varying gridworld environments (
|
||||
- [MiniGrid](https://minigrid.farama.org/environments/minigrid/), ~~- BabyAI~~
|
||||
- Crafter,
|
||||
- 2019 [Rogue-gym,](https://github.com/kngwyu/rogue-gym)
|
||||
- 2020 MarsExplorer, maxe exploration. 1M steps
|
||||
- NLE,
|
||||
- MiniHack;
|
||||
- [gym\_nethack](http://campbelljc.com/research/gym_nethack/)
|
||||
- 2018 [rogueinabox](https://github.com/rogueinabox/rogueinabox)
|
||||
- [rogue-gym](https://github.com/kngwyu/rogue-gym)
|
||||
- [MiniGrid](https://github.com/maximecb/gym-minigrid)
|
||||
- 2019 [CoinRun](https://github.com/openai/coinrun) no traction or maintanance
|
||||
- [MineRL](http://minerl.io/docs)
|
||||
- [Project Malmo](https://www.microsoft.com/en-us/research/project/project-malmo/) miencraft
|
||||
- [OpenAI Procgen Benchmark](https://openai.com/blog/procgen-benchmark/) 200M steps
|
||||
- 2020 [Obstacle Tower](https://github.com/Unity-Technologies/obstacle-tower-env) - 3d slow
|
||||
|
||||
non-PCG observation-varying continuous control environments
|
||||
- (RoboSuite, DMC-Remastered, DMC-GB, DCS, KitchenShift, NaturalEnvs MuJoCo; Fan
|
||||
et al., 2021; Grigsby & Qi, 2020; Hansen & Wang, 2021; Stone et al., 2021; Xing et al., 2021a;
|
||||
Zhang et al., 2018a), and multi-task continuous control benchmarks which could be adapted
|
||||
to ZSG (CausalWorld, RLBench, Meta-world; Ahmed et al., 2020; James et al., 2019a; Yu
|
||||
et al., 2019).
|
||||
|
||||
## investigating model slowness
|
||||
|
||||
So it's all just that using the transformer to imagine takes almost 0.1s. but it's run so many more times than during training. All my ideas to speed it up don't work.
|
||||
|
||||
- [x] eval. no grad
|
||||
- [x] remove the call for adapter, causal mask each time?
|
||||
- [ ] lower rank?
|
||||
|
||||
|
||||
Ok so it's all just the
|
||||
- rollout, controller by max block size. 10x
|
||||
- the fact that actor_critic can use a larger batch, therefore 4-8x more samples
|
||||
- for each one it imagines 2
|
||||
|
||||
@@ -36,8 +36,8 @@ class ImagineOutput:
|
||||
class ActorCritic(nn.Module):
|
||||
def __init__(self, act_vocab_size, use_original_obs: bool = False, lstm_dim = 16) -> None:
|
||||
super().__init__()
|
||||
shrink = 8
|
||||
s = 2
|
||||
shrink = 1
|
||||
s = 1
|
||||
self.use_original_obs = use_original_obs
|
||||
self.conv1 = nn.Conv2d(3, 32//s, 3, stride=1, padding=1)
|
||||
self.maxp1 = nn.MaxPool2d(2, 2)
|
||||
|
||||
@@ -49,7 +49,8 @@ class Transformer(nn.Module):
|
||||
# @torch.cuda.amp.autocast(dtype=torch.bfloat16)
|
||||
def forward(self, sequences: torch.Tensor, past_keys_values: Optional[KeysValues] = None) -> torch.Tensor:
|
||||
assert past_keys_values is None or len(past_keys_values) == self.config.num_layers
|
||||
with set_adapter(self.model, "dynamics"), disable_causal_mask(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
# with set_adapter(self.model, "dynamics"), disable_causal_mask(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
with torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
# sequences = sequences.to(torch.bfloat16)
|
||||
outputs = self.model(
|
||||
inputs_embeds=sequences,
|
||||
@@ -111,7 +112,9 @@ def load_pretrained_model(config, device="cuda:0"):
|
||||
)
|
||||
base_model_peft = peft.get_peft_model(base_model, peft_config)
|
||||
base_model_peft.add_adapter("dynamics", peft_config)
|
||||
base_model_peft.set_adapter("dynamics")
|
||||
print(base_model_peft.print_trainable_parameters())
|
||||
disable_causal_mask()
|
||||
return base_model_peft
|
||||
|
||||
@contextmanager
|
||||
|
||||
@@ -46,6 +46,7 @@ class Trainer:
|
||||
self.reconstructions_dir = self.media_dir / 'reconstructions'
|
||||
|
||||
if not cfg.common.resume:
|
||||
print('cwd', Path.cwd())
|
||||
config_dir = Path('config')
|
||||
config_path = config_dir / 'trainer.yaml'
|
||||
config_dir.mkdir(exist_ok=False, parents=False)
|
||||
|
||||
Reference in New Issue
Block a user