MCPcopy Create free account
hub / github.com/WuTao-CS/CustomCrafter / print_learnable_params

Function print_learnable_params

custom.py:29–55  ·  view source on GitHub ↗
(model)

Source from the content-addressed store, hash-verified

27 return sorted(k for k in vars(args) if getattr(opt, k) != getattr(args, k))
28
29def print_learnable_params(model):
30 # 获取优化器
31 optimizer = model.configure_optimizers()
32 cnt = 0
33 print("Learnable parameters: ")
34 for name, param in model.named_parameters():
35 if param.requires_grad:
36 print(name)
37 cnt+=1
38 print("Total number of learnable parameters: ", cnt)
39 # 如果有多个优化器,你可能需要遍历它们
40 if isinstance(optimizer, list):
41 for idx, opt in enumerate(optimizer):
42 print(f"Optimizer {idx}:")
43 for param_group in opt.param_groups:
44 for name, param in model.lightning_module.named_parameters():
45 if param.requires_grad:
46 print(f"Parameter: {name}, Size: {param.size()}")
47 else:
48 cnt_opt = 0
49 for param_group in optimizer.param_groups:
50 for item in param_group['params']:
51 # print(item.size())
52 cnt_opt+=1
53 print("Total number of learnable parameters in optimizer: ", cnt_opt)
54 # print(param_group['params'][0].size())
55 return
56
57def load_model_from_config(config, ckpt, verbose=False):
58 print(f"Loading model from {ckpt}")

Callers 1

custom.pyFile · 0.85

Calls 1

configure_optimizersMethod · 0.45

Tested by

no test coverage detected