File size: 1,990 Bytes
c881b77
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
import torch
import itertools

def get_optimizer(config, model):
    if config.training.optimizer.name == "AdamW":
        return torch.optim.AdamW(
            itertools.chain(model.unet.parameters(), model.image_proj_model.parameters()),
            lr=config.training.optimizer.learning_rate,
            betas=config.training.optimizer.adam_beta,
            weight_decay=config.training.optimizer.weight_decay,
            eps=config.training.optimizer.adam_epsilon,
        )
    if config.training.optimizer.name == "AdamWGating":
        base_lr = config.training.optimizer.learning_rate
        gating_multiplier = 100.0 

        gating_params, regular_params = [], []

        for module in (model.unet, model.image_proj_model):
            for name, param in module.named_parameters():
                if not param.requires_grad:
                    continue
                if "gating_param" in name:
                    gating_params.append(param)
                else:
                    regular_params.append(param)

        param_groups = [
            {
                'params': regular_params, 
                'lr': base_lr
            },
            {
                'params': gating_params, 
                'lr': base_lr * gating_multiplier
            }
        ]
        
        # if len(gating_params) > 0:
        #     print(f"[Optimizer] Found {len(gating_params)} gating parameters. Setting LR to {base_lr * gating_multiplier} (x{gating_multiplier})")
        # else:
        #     print("[Optimizer] Warning: No 'gating_param' found in trainable parameters!")

        return torch.optim.AdamW(
            param_groups,
            lr=base_lr,
            betas=config.training.optimizer.adam_beta,
            weight_decay=config.training.optimizer.weight_decay,
            eps=config.training.optimizer.adam_epsilon,
        )
    else:
        raise ValueError(f"Unsupported optimizer: {config.training.optimizer.name}. Supported optimizers: AdamW.")