This commit is contained in:
wassname
2023-11-19 08:17:38 +08:00
parent 9b399031ca
commit 36ea66d0ae
8 changed files with 79 additions and 38 deletions
+1 -12
View File
@@ -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
View File
@@ -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
+36
View File
@@ -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
View File
@@ -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,
}
+1 -1
View File
@@ -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)
+12 -5
View File
@@ -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
+10 -1
View File
@@ -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
View File
@@ -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)