mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-08-22 12:10:08 +08:00
33 KiB
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')
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 [ ]:
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
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)
[0;31m---------------------------------------------------------------------------[0m [0;31mAssertionError[0m Traceback (most recent call last) [1;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/notebooks/01_debug_models.ipynb Cell 24[0m line [0;36m1 [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>[0m get_ipython()[39m.[39;49mrun_cell_magic([39m'[39;49m[39mtimeit[39;49m[39m'[39;49m, [39m'[39;49m[39m'[39;49m, [39m"[39;49m[39m# ofc it[39;49m[39m'[39;49m[39ms the transformer that[39;49m[39m'[39;49m[39ms slow. I guess we just call it was more than during training[39;49m[39m\n[39;49;00m[39mpast_keys_values = wm_env.keys_values_wm[39;49m[39m\n[39;49;00m[39mx = world_model.transformer(sequences, past_keys_values)[39;49m[39m\n[39;49;00m[39m"[39;49m) File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/interactiveshell.py:2515[0m, in [0;36mInteractiveShell.run_cell_magic[0;34m(self, magic_name, line, cell)[0m [1;32m 2513[0m [39mwith[39;00m [39mself[39m[39m.[39mbuiltin_trap: [1;32m 2514[0m args [39m=[39m (magic_arg_s, cell) [0;32m-> 2515[0m result [39m=[39m fn([39m*[39;49margs, [39m*[39;49m[39m*[39;49mkwargs) [1;32m 2517[0m [39m# The code below prevents the output from being displayed[39;00m [1;32m 2518[0m [39m# when using magics with decorator @output_can_be_silenced[39;00m [1;32m 2519[0m [39m# when the last Python token in the expression is a ';'.[39;00m [1;32m 2520[0m [39mif[39;00m [39mgetattr[39m(fn, magic[39m.[39mMAGIC_OUTPUT_CAN_BE_SILENCED, [39mFalse[39;00m): File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/magics/execution.py:1189[0m, in [0;36mExecutionMagics.timeit[0;34m(self, line, cell, local_ns)[0m [1;32m 1186[0m [39mif[39;00m time_number [39m>[39m[39m=[39m [39m0.2[39m: [1;32m 1187[0m [39mbreak[39;00m [0;32m-> 1189[0m all_runs [39m=[39m timer[39m.[39;49mrepeat(repeat, number) [1;32m 1190[0m best [39m=[39m [39mmin[39m(all_runs) [39m/[39m number [1;32m 1191[0m worst [39m=[39m [39mmax[39m(all_runs) [39m/[39m number File [0;32m~/miniforge3/lib/python3.9/timeit.py:205[0m, in [0;36mTimer.repeat[0;34m(self, repeat, number)[0m [1;32m 203[0m r [39m=[39m [] [1;32m 204[0m [39mfor[39;00m i [39min[39;00m [39mrange[39m(repeat): [0;32m--> 205[0m t [39m=[39m [39mself[39;49m[39m.[39;49mtimeit(number) [1;32m 206[0m r[39m.[39mappend(t) [1;32m 207[0m [39mreturn[39;00m r File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/IPython/core/magics/execution.py:173[0m, in [0;36mTimer.timeit[0;34m(self, number)[0m [1;32m 171[0m gc[39m.[39mdisable() [1;32m 172[0m [39mtry[39;00m: [0;32m--> 173[0m timing [39m=[39m [39mself[39;49m[39m.[39;49minner(it, [39mself[39;49m[39m.[39;49mtimer) [1;32m 174[0m [39mfinally[39;00m: [1;32m 175[0m [39mif[39;00m gcold: File [0;32m<magic-timeit>:3[0m, in [0;36minner[0;34m(_it, _timer)[0m File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/nn/modules/module.py:1518[0m, in [0;36mModule._wrapped_call_impl[0;34m(self, *args, **kwargs)[0m [1;32m 1516[0m [39mreturn[39;00m [39mself[39m[39m.[39m_compiled_call_impl([39m*[39margs, [39m*[39m[39m*[39mkwargs) [39m# type: ignore[misc][39;00m [1;32m 1517[0m [39melse[39;00m: [0;32m-> 1518[0m [39mreturn[39;00m [39mself[39;49m[39m.[39;49m_call_impl([39m*[39;49margs, [39m*[39;49m[39m*[39;49mkwargs) File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/nn/modules/module.py:1527[0m, in [0;36mModule._call_impl[0;34m(self, *args, **kwargs)[0m [1;32m 1522[0m [39m# If we don't have any hooks, we want to skip the rest of the logic in[39;00m [1;32m 1523[0m [39m# this function, and just call forward.[39;00m [1;32m 1524[0m [39mif[39;00m [39mnot[39;00m ([39mself[39m[39m.[39m_backward_hooks [39mor[39;00m [39mself[39m[39m.[39m_backward_pre_hooks [39mor[39;00m [39mself[39m[39m.[39m_forward_hooks [39mor[39;00m [39mself[39m[39m.[39m_forward_pre_hooks [1;32m 1525[0m [39mor[39;00m _global_backward_pre_hooks [39mor[39;00m _global_backward_hooks [1;32m 1526[0m [39mor[39;00m _global_forward_hooks [39mor[39;00m _global_forward_pre_hooks): [0;32m-> 1527[0m [39mreturn[39;00m forward_call([39m*[39;49margs, [39m*[39;49m[39m*[39;49mkwargs) [1;32m 1529[0m [39mtry[39;00m: [1;32m 1530[0m result [39m=[39m [39mNone[39;00m File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/transformer.py:69[0m, in [0;36mTransformer.forward[0;34m(self, sequences, past_keys_values)[0m [1;32m 66[0m [39m# k_size = (x.shape[0], x.shape[1], x.shape[1], 1)[39;00m [1;32m 67[0m [39m# v_size = past_keys_values[0]._v_cache._cache.size()[39;00m [1;32m 68[0m v_size [39m=[39m (k_size[[39m0[39m], k_size[[39m1[39m], x[39m.[39mshape[[39m1[39m], k_size[[39m3[39m]) [0;32m---> 69[0m past_keys_values[[39m0[39;49m][39m.[39;49mupdate(torch[39m.[39;49mrand(v_size), torch[39m.[39;49mrand(v_size)) [1;32m 70[0m [39mreturn[39;00m x File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/kv_caching.py:59[0m, in [0;36mKVCache.update[0;34m(self, k, v)[0m [1;32m 58[0m [39mdef[39;00m [39mupdate[39m([39mself[39m, k: torch[39m.[39mTensor, v: torch[39m.[39mTensor): [0;32m---> 59[0m [39mself[39;49m[39m.[39;49m_k_cache[39m.[39;49mupdate(k) [1;32m 60[0m [39mself[39m[39m.[39m_v_cache[39m.[39mupdate(v) File [0;32m/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/kv_caching.py:33[0m, in [0;36mCache.update[0;34m(self, x)[0m [1;32m 31[0m [39mdef[39;00m [39mupdate[39m([39mself[39m, x: torch[39m.[39mTensor) [39m-[39m[39m>[39m [39mNone[39;00m: [1;32m 32[0m [39massert[39;00m (x[39m.[39mndim [39m==[39m [39mself[39m[39m.[39m_cache[39m.[39mndim) [39mand[39;00m [39mall[39m([x[39m.[39msize(i) [39m==[39m [39mself[39m[39m.[39m_cache[39m.[39msize(i) [39mfor[39;00m i [39min[39;00m ([39m0[39m, [39m1[39m, [39m3[39m)]) [0;32m---> 33[0m [39massert[39;00m [39mself[39m[39m.[39m_size [39m+[39m x[39m.[39msize([39m2[39m) [39m<[39m[39m=[39m [39mself[39m[39m.[39m_cache[39m.[39mshape[[39m2[39m] [1;32m 34[0m [39mself[39m[39m.[39m_cache [39m=[39m AssignWithoutInplaceCheck[39m.[39mapply([39mself[39m[39m.[39m_cache, x, [39m2[39m, [39mself[39m[39m.[39m_size, [39mself[39m[39m.[39m_size [39m+[39m x[39m.[39msize([39m2[39m)) [1;32m 35[0m [39mself[39m[39m.[39m_size [39m+[39m[39m=[39m x[39m.[39msize([39m2[39m) [0;31mAssertionError[0m:
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)
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)
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 [ ]: