diff --git a/.vscode/launch.json b/.vscode/launch.json index 9044a8c..dbe115f 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -14,7 +14,7 @@ "autoReload": {"enable": true,}, "env": {"WANDB_MODE":"disabled"}, "args": [ - "'wandb.mode=disabled", + // "'wandb.mode=disabled", // "env.train.id=BreakoutNoFrameskip-v4", "env.train.id=CrafterReward-v1", // # make it start early @@ -34,19 +34,8 @@ "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", ] } ] diff --git a/config/trainer.yaml b/config/trainer.yaml index 390a3ac..b4fe7c5 100644 --- a/config/trainer.yaml +++ b/config/trainer.yaml @@ -60,7 +60,7 @@ training: start_after_epochs: 5 steps_per_epoch: 200 world_model: - batch_num_samples: 8 # pretrained models use lots of + batch_num_samples: 8 # pretrained models use lots of ram grad_acc_steps: 1 max_grad_norm: 10.0 weight_decay: 0.01 @@ -70,7 +70,7 @@ training: batch_num_samples: 16 grad_acc_steps: 1 max_grad_norm: 10.0 - start_after_epochs: 50 + start_after_epochs: 1500 steps_per_epoch: 40 imagine_horizon: ${common.sequence_length} burn_in: 20 diff --git a/research_journal.md b/research_journal.md index 0d40317..b7e2a98 100644 --- a/research_journal.md +++ b/research_journal.md @@ -318,3 +318,39 @@ Ok so it's all just the # 2023-11-18 06:17:55 It trained overnight, now I would like to view a replay + +Hmm "delta-IRIS" ∆-IRIS +https://openreview.net/forum?id=o8IDoZggqO +∆-IRIS encodes +new frames by attending to the ongoing trajectory, effec- +tively describing deltas between timesteps. +This new ap- +proach drastically reduces the number of tokens to encode +frames, since they are not encoded independently as in IRIS. +In the Crafter benchmark (Hafner, 2022), ∆-IRIS unlocks +16 out of 22 objectives at the 10M frames mark + + +# 2023-11-18 16:16:16 + +Why is it not learning? It's because the dynamics model is total BS!!! + +- [ ] Well lets try training it for longer then. It's cheap to train so.. +- [ ] also maybe train tokenizer and model together? I have a lot of frozen layers, including the embeddings... so might be better + - [ ] oh no we do have an unforzen embedder before the transformer or more layers + - maybe I need a higher rank lora? after all I'm changing a lot from text tokens + - maybe no tokens, bypass to embedder? + + +# 2023-11-19 06:45:34 + +So I tried just trainign the world model for 200 epochs. And with a post_embedding layer. It helped the flickering. But not enougth to actually go for the obvious local minima of the next state equals the last + + + + +idea +- bypass embedding?, but wait dreamerv3 needed quant z... + - yes I am bypassing it by passing in the input_embeds... but maybe I shouldn't +- use same embedding everywhere. e.g. model embedding in encoder decoder? + - Our embedings is (embed_tokens): Embedding(32000, 2048). So we would need to encode to 32000! diff --git a/src/game/keymap.py b/src/game/keymap.py index 735eb8c..e912868 100644 --- a/src/game/keymap.py +++ b/src/game/keymap.py @@ -107,22 +107,22 @@ EMPTY_KEYMAP = { } CRAFTER_KEYMAP = { - pygame.K_a: 'move_left', - pygame.K_d: 'move_right', - pygame.K_w: 'move_up', - pygame.K_s: 'move_down', - pygame.K_SPACE: 'do', - pygame.K_TAB: 'sleep', + pygame.K_a: 1, + pygame.K_d: 2, + pygame.K_w: 3, + pygame.K_s: 4, + pygame.K_SPACE: 5, + pygame.K_TAB: 6, - pygame.K_r: 'place_stone', - pygame.K_t: 'place_table', - pygame.K_f: 'place_furnace', - pygame.K_p: 'place_plant', + pygame.K_r: 7, + pygame.K_t: 8, + pygame.K_f: 9, + pygame.K_p: 10, - pygame.K_1: 'make_wood_pickaxe', - pygame.K_2: 'make_stone_pickaxe', - pygame.K_3: 'make_iron_pickaxe', - pygame.K_4: 'make_wood_sword', - pygame.K_5: 'make_stone_sword', - pygame.K_6: 'make_iron_sword', + pygame.K_1: 11, + pygame.K_2: 12, + pygame.K_3: 13, + pygame.K_4: 14, + pygame.K_5: 15, + pygame.K_6: 16, } diff --git a/src/models/tokenizer/tokenizer.py b/src/models/tokenizer/tokenizer.py index 1b8723d..e69c18e 100644 --- a/src/models/tokenizer/tokenizer.py +++ b/src/models/tokenizer/tokenizer.py @@ -28,7 +28,7 @@ class Tokenizer(nn.Module): self.vocab_size = vocab_size self.encoder = encoder self.pre_quant_conv = torch.nn.Conv2d(encoder.config.z_channels, embed_dim, 1) - self.embedding = nn.Embedding(vocab_size, embed_dim) + self.embedding = nn.Embedding(vocab_size, embed_dim) # TODO: use model embed? self.post_quant_conv = torch.nn.Conv2d(embed_dim, decoder.config.z_channels, 1) self.decoder = decoder self.embedding.weight.data.uniform_(-1.0 / vocab_size, 1.0 / vocab_size) diff --git a/src/models/transformer.py b/src/models/transformer.py index 5ccdc70..89a2b25 100644 --- a/src/models/transformer.py +++ b/src/models/transformer.py @@ -46,18 +46,15 @@ class Transformer(nn.Module): device = self.ln_f.weight.device # Assumption that all submodules are on the same device return KeysValues(n, self.config.num_heads, max_tokens, self.config.embed_dim, self.config.num_layers, device) - # @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 torch.cuda.amp.autocast(dtype=torch.bfloat16): - # sequences = sequences.to(torch.bfloat16) outputs = self.model( inputs_embeds=sequences, return_dict=True, output_hidden_states=True, ) - x = outputs.logits#.to(torch.float32) + x = outputs.logits x = self.ln_f(x) # fake it, since it's used to keep track of steps @@ -114,7 +111,7 @@ def load_pretrained_model(config, device="cuda:0"): base_model_peft.add_adapter("dynamics", peft_config) base_model_peft.set_adapter("dynamics") print(base_model_peft.print_trainable_parameters()) - disable_causal_mask() + disable_causal_mask_always() return base_model_peft @contextmanager @@ -130,6 +127,16 @@ def set_adapter(model, adapter_name): finally: model.set_adapter(old_adapter_name) +def disable_causal_mask_always(): + import transformers.models.llama.modeling_llama as modeling + + decoder_fn = modeling._make_causal_mask + + def encoder_fn(*args, **kwargs): + return torch.zeros_like(decoder_fn(*args, **kwargs)) + + modeling._make_causal_mask = encoder_fn + @contextmanager def disable_causal_mask(): import transformers.models.llama.modeling_llama as modeling diff --git a/src/models/world_model.py b/src/models/world_model.py index eecfe60..723b733 100644 --- a/src/models/world_model.py +++ b/src/models/world_model.py @@ -41,6 +41,13 @@ class WorldModel(nn.Module): block_masks=[act_tokens_pattern, obs_tokens_pattern], embedding_tables=nn.ModuleList([nn.Embedding(act_vocab_size, config.embed_dim), nn.Embedding(obs_vocab_size, config.embed_dim)]) ) + self.post_embed = nn.Sequential( + nn.Linear(config.embed_dim, config.embed_dim), + nn.ReLU(), + nn.Linear(config.embed_dim, config.embed_dim), + nn.ReLU(), + nn.Linear(config.embed_dim, config.embed_dim) + ) self.head_observations = Head( max_blocks=config.max_blocks, @@ -86,7 +93,8 @@ class WorldModel(nn.Module): prev_steps = 0 if past_keys_values is None else past_keys_values.size sequences = self.embedder(tokens, num_steps, prev_steps) + self.pos_emb(prev_steps + torch.arange(num_steps, device=tokens.device)) - + # [batch=8, num_steps=170, embed_size=2048] + sequences = self.post_embed(sequences) x = self.transformer(sequences, past_keys_values) logits_observations = self.head_observations(x, num_steps=num_steps, prev_steps=prev_steps) @@ -98,6 +106,7 @@ class WorldModel(nn.Module): def compute_loss(self, batch: Batch, tokenizer: Tokenizer, **kwargs: Any) -> LossWithIntermediateLosses: with torch.no_grad(): + # [B=8, S=10, Colors=3, H=64, W=64] -> [B=8, S=10, 16] obs_tokens = tokenizer.encode(batch['observations'], should_preprocess=True).tokens # (BL, K) act_tokens = rearrange(batch['actions'], 'b l -> b l 1') diff --git a/src/trainer.py b/src/trainer.py index 9402c60..af0460c 100644 --- a/src/trainer.py +++ b/src/trainer.py @@ -137,11 +137,11 @@ class Trainer: 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() + self.agent.tokenizer.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)