MCPcopy Create free account
hub / github.com/adobe-research/custom-diffusion / compress

Function compress

src/compress.py:8–54  ·  view source on GitHub ↗
(delta_ckpt, ckpt, diffuser=False, compression_ratio=0.6, device='cuda')

Source from the content-addressed store, hash-verified

6
7
8def 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
57def parse_args():

Callers 1

compress.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected