MCPcopy Create free account
hub / github.com/YuxuanSnow/Human3Diffusion / print_model_info

Function print_model_info

train_MultiviewDiffusion_diffusion.py:278–284  ·  view source on GitHub ↗
(model)

Source from the content-addressed store, hash-verified

276 )
277
278 def print_model_info(model):
279 print("="*20)
280 print("model name: ", type(model).__name__)
281 print("learnable parameters(M): ", sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6)
282 print("non-learnable parameters(M): ", sum(p.numel() for p in model.parameters() if not p.requires_grad) / 1e6)
283 print("total parameters(M): ", sum(p.numel() for p in model.parameters()) / 1e6)
284 print("model size(MB): ", sum(p.numel() * p.element_size() for p in model.parameters()) / 1024 / 1024)
285
286 print_model_info(unet)
287 print_model_info(vae)

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected