MCPcopy Create free account
hub / github.com/CompVis/zigma / initialize_train_state

Function initialize_train_state

utils/train_state_utils.py:71–89  ·  view source on GitHub ↗
(config, model, model_ema , device)

Source from the content-addressed store, hash-verified

69
70
71def 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

Callers

nothing calls this directly

Calls 4

ema_updateMethod · 0.95
toMethod · 0.95
cnt_paramsFunction · 0.85
TrainStateClass · 0.85

Tested by

no test coverage detected