Fix bugs and improve performance in training

process
This commit is contained in:
wassname
2023-11-13 07:00:17 +08:00
parent c8f59b0504
commit b97a499a19
5 changed files with 55 additions and 9 deletions
+5 -5
View File
@@ -17,11 +17,11 @@
"env.train.id=BreakoutNoFrameskip-v4",
// # make it start early
"training.tokenizer.start_after_epochs=1",
"training.world_model.start_after_epochs=2",
"training.actor_critic.start_after_epochs=3",
"training.tokenizer.steps_per_epoch=40",
"training.world_model.steps_per_epoch=40",
"training.actor_critic.steps_per_epoch=40",
"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",
]
}
]
+36
View File
@@ -111,3 +111,39 @@ hmm it still happens in the original repo with my debug params. maybe it's my de
... trying a full run without my debug params...
note trains.world_model.batch_num_samples:4 fill 20GB gpu ram for the 3b stability ai llm
ok even with a full run I get the error. I think it's a bug in the original repo. I'll try to debug it there.
Epoch 51 / 600
Experience collection (train_dataset): 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [00:03<00:00, 59.91it/s]
Training tokenizer: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [00:17<00:00, 11.53it/s]
Training world_model: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [02:11<00:00, 1.53it/s]
Training actor_critic: 0%| | 0/200 [00:00<?, ?it/s]
Error executing job with overrides: ['env.train.id=BreakoutNoFrameskip-v4', 'common.device=cuda:0', 'wandb.mode=offline']
Traceback (most recent call last):
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/main.py", line 10, in main
trainer.run()
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/trainer.py", line 111, in run
to_log += self.train_agent(epoch)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/trainer.py", line 146, in train_agent
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)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/trainer.py", line 161, in train_component
losses = component.compute_loss(batch, **kwargs_loss) / grad_acc_steps
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/actor_critic.py", line 102, in compute_loss
outputs = self.imagine(batch, tokenizer, world_model, horizon=imagine_horizon)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/models/actor_critic.py", line 149, in imagine
obs, reward, done, _ = wm_env.step(action_token, should_predict_next_obs=(k < horizon - 1))
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/utils/_contextlib.py", line 115, in decorate_context
return func(*args, **kwargs)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/src/envs/world_model_env.py", line 75, in step
reward = Categorical(logits=outputs_wm.logits_rewards).sample().float().cpu().numpy().reshape(-1) - 1 # (B,)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/distributions/categorical.py", line 70, in __init__
super().__init__(batch_shape, validate_args=validate_args)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/distributions/distribution.py", line 66, in __init__
valid = constraint.check(value)
File "/media/wassname/SGIronWolf/projects5/worldmodels/iris_bigvae/.venv/lib/python3.9/site-packages/torch/distributions/constraints.py", line 226, in check
result = result.reshape(
RuntimeError: cannot reshape tensor of 0 elements into shape [8, 0, -1] because the unspecified dimension size -1 can be any value and is ambiguous
Oh maybe it's because we don't keep track of KV cache, but it's actually used to track number of steps!!
+3 -2
View File
@@ -56,7 +56,7 @@ class WorldModelEnv:
def step(self, action: Union[int, np.ndarray, torch.LongTensor], should_predict_next_obs: bool = True) -> None:
assert self.keys_values_wm is not None and self.num_observations_tokens is not None
num_passes = 1 + self.num_observations_tokens if should_predict_next_obs else 1
num_passes = 2 + self.num_observations_tokens if should_predict_next_obs else 1
output_sequence, obs_tokens = [], []
@@ -71,7 +71,8 @@ class WorldModelEnv:
outputs_wm = self.world_model(token, past_keys_values=self.keys_values_wm)
output_sequence.append(outputs_wm.output_sequence)
if k == 0:
# if k == 0:
if self.world_model(token, past_keys_values=self.keys_values_wm).logits_rewards.shape[1] > 0:
reward = Categorical(logits=outputs_wm.logits_rewards).sample().float().cpu().numpy().reshape(-1) - 1 # (B,)
done = Categorical(logits=outputs_wm.logits_ends).sample().cpu().numpy().astype(bool).reshape(-1) # (B,)
+2 -1
View File
@@ -31,7 +31,8 @@ class Head(Slicer):
self.head_module = head_module
def forward(self, x: torch.Tensor, num_steps: int, prev_steps: int) -> torch.Tensor:
x_sliced = x[:, self.compute_slice(num_steps, prev_steps)] # x is (B, T, E)
s = self.compute_slice(num_steps, prev_steps)
x_sliced = x[:, s] # x is (B, T, E)
return self.head_module(x_sliced)
+9 -1
View File
@@ -117,7 +117,7 @@ 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) == len(self.blocks)
assert past_keys_values is None or len(past_keys_values) == self.config.num_layers
sequences = sequences.to(torch.bfloat16)
outputs = self.model(
inputs_embeds=sequences,
@@ -126,6 +126,14 @@ class Transformer(nn.Module):
)
x = outputs.logits.to(torch.float32)
x = self.ln_f(x)
# fake it, since it's used to keep track of steps
if past_keys_values is not None:
k_size = past_keys_values[0]._k_cache._cache.size()
k_size = (*k_size[:2], 1, *k_size[3:])
v_size = past_keys_values[0]._v_cache._cache.size()
v_size = (*v_size[:2], 1, *v_size[3:])
past_keys_values[0].update(torch.rand(k_size), torch.rand(v_size))
return x