(model)
| 27 | return sorted(k for k in vars(args) if getattr(opt, k) != getattr(args, k)) |
| 28 | |
| 29 | def 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 | |
| 57 | def load_model_from_config(config, ckpt, verbose=False): |
| 58 | print(f"Loading model from {ckpt}") |
no test coverage detected