(config, model, model_ema , device)
| 69 | |
| 70 | |
| 71 | def initialize_train_state(config, model, model_ema , device): |
| 72 | params = [] |
| 73 | params += model.parameters() |
| 74 | model_ema.eval() |
| 75 | logging.warning(f"nnet has {cnt_params(model)} parameters") |
| 76 | optimizer = torch.optim.AdamW( |
| 77 | model.parameters(), lr=config.optim.lr, weight_decay=config.optim.wd |
| 78 | ) |
| 79 | |
| 80 | train_state = TrainState( |
| 81 | optimizer=optimizer, |
| 82 | step=0, |
| 83 | model=model, |
| 84 | model_ema=model_ema, |
| 85 | ) |
| 86 | train_state.ema_update(0) |
| 87 | if device is not None: |
| 88 | train_state.to(device) |
| 89 | return train_state |
nothing calls this directly
no test coverage detected