| 6 | |
| 7 | |
| 8 | def compress(delta_ckpt, ckpt, diffuser=False, compression_ratio=0.6, device='cuda'): |
| 9 | st = torch.load(f'{delta_ckpt}') |
| 10 | |
| 11 | if not diffuser: |
| 12 | compressed_key = 'state_dict' |
| 13 | compressed_st = {compressed_key: {}} |
| 14 | pretrained_st = torch.load(ckpt)['state_dict'] |
| 15 | if 'embed' in st['state_dict']: |
| 16 | compressed_st['state_dict']['embed'] = st['state_dict']['embed'] |
| 17 | del st['state_dict']['embed'] |
| 18 | |
| 19 | st = st['state_dict'] |
| 20 | else: |
| 21 | from diffusers import StableDiffusionPipeline |
| 22 | compressed_key = 'unet' |
| 23 | compressed_st = {compressed_key: {}} |
| 24 | pretrained_st = StableDiffusionPipeline.from_pretrained(ckpt, torch_dtype=torch.float16).to("cuda") |
| 25 | pretrained_st = pretrained_st.unet.state_dict() |
| 26 | if 'modifier_token' in st: |
| 27 | compressed_st['modifier_token'] = st['modifier_token'] |
| 28 | st = st['unet'] |
| 29 | |
| 30 | print("getting compression") |
| 31 | layers = list(st.keys()) |
| 32 | for name in layers: |
| 33 | if 'to_k' in name or 'to_v' in name: |
| 34 | W = st[name].to(device) |
| 35 | Wpretrain = pretrained_st[name].clone().to(device) |
| 36 | deltaW = W-Wpretrain |
| 37 | |
| 38 | u, s, vt = torch.linalg.svd(deltaW.clone()) |
| 39 | |
| 40 | explain = 0 |
| 41 | all_ = (s).sum() |
| 42 | for i, t in enumerate(s): |
| 43 | explain += t/(all_) |
| 44 | if explain > compression_ratio: |
| 45 | break |
| 46 | |
| 47 | compressed_st[compressed_key][f'{name}'] = {} |
| 48 | compressed_st[compressed_key][f'{name}']['u'] = (u[:, :i]@torch.diag(s)[:i, :i]).clone() |
| 49 | compressed_st[compressed_key][f'{name}']['v'] = vt[:i].clone() |
| 50 | else: |
| 51 | compressed_st[compressed_key][f'{name}'] = st[name] |
| 52 | |
| 53 | name = delta_ckpt.replace('delta', 'compressed_delta') |
| 54 | torch.save(compressed_st, f'{name}') |
| 55 | |
| 56 | |
| 57 | def parse_args(): |