| 38 | |
| 39 | |
| 40 | def initialize_unet(rank_pix, rank_sem, return_lora_module_names=False, pretrained_model_path=None): |
| 41 | unet = UNet2DConditionModel.from_pretrained(pretrained_model_path, subfolder="unet") |
| 42 | unet.requires_grad_(False) |
| 43 | unet.train() |
| 44 | |
| 45 | l_target_modules_encoder_pix, l_target_modules_decoder_pix, l_modules_others_pix = [], [], [] |
| 46 | l_target_modules_encoder_sem, l_target_modules_decoder_sem, l_modules_others_sem = [], [], [] |
| 47 | l_grep = ["to_k", "to_q", "to_v", "to_out.0", "conv", "conv1", "conv2", "conv_in", "conv_shortcut", "conv_out", "proj_out", "proj_in", "ff.net.2", "ff.net.0.proj"] |
| 48 | for n, p in unet.named_parameters(): |
| 49 | check_flag = 0 |
| 50 | if "bias" in n or "norm" in n: |
| 51 | continue |
| 52 | for pattern in l_grep: |
| 53 | if pattern in n and ("down_blocks" in n or "conv_in" in n): |
| 54 | l_target_modules_encoder_pix.append(n.replace(".weight","")) |
| 55 | l_target_modules_encoder_sem.append(n.replace(".weight","")) |
| 56 | break |
| 57 | elif pattern in n and ("up_blocks" in n or "conv_out" in n): |
| 58 | l_target_modules_decoder_pix.append(n.replace(".weight","")) |
| 59 | l_target_modules_decoder_sem.append(n.replace(".weight","")) |
| 60 | break |
| 61 | elif pattern in n: |
| 62 | l_modules_others_pix.append(n.replace(".weight","")) |
| 63 | l_modules_others_sem.append(n.replace(".weight","")) |
| 64 | break |
| 65 | |
| 66 | lora_conf_encoder_pix = LoraConfig(r=rank_pix, init_lora_weights="gaussian",target_modules=l_target_modules_encoder_pix) |
| 67 | lora_conf_decoder_pix = LoraConfig(r=rank_pix, init_lora_weights="gaussian",target_modules=l_target_modules_decoder_pix) |
| 68 | lora_conf_others_pix = LoraConfig(r=rank_pix, init_lora_weights="gaussian",target_modules=l_modules_others_pix) |
| 69 | lora_conf_encoder_sem = LoraConfig(r=rank_sem, init_lora_weights="gaussian",target_modules=l_target_modules_encoder_sem) |
| 70 | lora_conf_decoder_sem = LoraConfig(r=rank_sem, init_lora_weights="gaussian",target_modules=l_target_modules_decoder_sem) |
| 71 | lora_conf_others_sem = LoraConfig(r=rank_sem, init_lora_weights="gaussian",target_modules=l_modules_others_sem) |
| 72 | |
| 73 | unet.add_adapter(lora_conf_encoder_pix, adapter_name="default_encoder_pix") |
| 74 | unet.add_adapter(lora_conf_decoder_pix, adapter_name="default_decoder_pix") |
| 75 | unet.add_adapter(lora_conf_others_pix, adapter_name="default_others_pix") |
| 76 | unet.add_adapter(lora_conf_encoder_sem, adapter_name="default_encoder_sem") |
| 77 | unet.add_adapter(lora_conf_decoder_sem, adapter_name="default_decoder_sem") |
| 78 | unet.add_adapter(lora_conf_others_sem, adapter_name="default_others_sem") |
| 79 | |
| 80 | if return_lora_module_names: |
| 81 | return unet, l_target_modules_encoder_pix, l_target_modules_decoder_pix, l_modules_others_pix, l_target_modules_encoder_sem, l_target_modules_decoder_sem, l_modules_others_sem |
| 82 | else: |
| 83 | return unet |
| 84 | |
| 85 | |
| 86 | class CSDLoss(torch.nn.Module): |