mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-10-02 12:20:43 +08:00
debugging model size and speed
This commit is contained in:
1 parent
a778296e40
commit
1dabd79251
15 files changed
+1088
-15
No files matched your search
@@ -36,8 +36,8 @@ class ImagineOutput:
|
||||
class ActorCritic(nn.Module):
|
||||
def __init__(self, act_vocab_size, use_original_obs: bool = False, lstm_dim = 16) -> None:
|
||||
super().__init__()
|
||||
shrink = 8
|
||||
s = 2
|
||||
shrink = 1
|
||||
s = 1
|
||||
self.use_original_obs = use_original_obs
|
||||
self.conv1 = nn.Conv2d(3, 32//s, 3, stride=1, padding=1)
|
||||
self.maxp1 = nn.MaxPool2d(2, 2)
|
||||
|
||||
@@ -49,7 +49,8 @@ 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) == self.config.num_layers
|
||||
with set_adapter(self.model, "dynamics"), disable_causal_mask(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
# 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,
|
||||
@@ -111,7 +112,9 @@ def load_pretrained_model(config, device="cuda:0"):
|
||||
)
|
||||
base_model_peft = peft.get_peft_model(base_model, peft_config)
|
||||
base_model_peft.add_adapter("dynamics", peft_config)
|
||||
base_model_peft.set_adapter("dynamics")
|
||||
print(base_model_peft.print_trainable_parameters())
|
||||
disable_causal_mask()
|
||||
return base_model_peft
|
||||
|
||||
@contextmanager
|
||||
|
||||
@@ -46,6 +46,7 @@ class Trainer:
|
||||
self.reconstructions_dir = self.media_dir / 'reconstructions'
|
||||
|
||||
if not cfg.common.resume:
|
||||
print('cwd', Path.cwd())
|
||||
config_dir = Path('config')
|
||||
config_path = config_dir / 'trainer.yaml'
|
||||
config_dir.mkdir(exist_ok=False, parents=False)
|
||||
|
||||
Reference in new issue
Block a user