mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-09 11:24:31 +08:00
misc
This commit is contained in:
Vendored
+1
-12
@@ -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",
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
+2
-2
@@ -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
|
||||
|
||||
@@ -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!
|
||||
|
||||
+16
-16
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
+1
-1
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user