MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / copy_layer_norm

Function copy_layer_norm

SwissArmyTransformer/examples/eva2/transform_param.py:90–116  ·  view source on GitHub ↗
(src, dst)

Source from the content-addressed store, hash-verified

88 assert (dst_dic[k].data == src_dic[k].data).all()
89
90def copy_layer_norm(src, dst):
91 src_ln = []
92 for k, v in src.named_parameters():
93 if 'norm' in k.lower():
94 src_ln.append((k, v))
95 dst_ln = []
96 for k, v in dst.named_parameters():
97 if 'layernorm' in k.lower():
98 dst_ln.append((k, v))
99 assert len(src_ln) == len(dst_ln)
100 for kvs, kvd in zip(src_ln, dst_ln):
101 assert kvd[1].data.shape == kvs[1].data.shape
102 kvd[1].data = kvs[1].data
103 assert (kvd[1].data == kvs[1].data).all()
104 src_ln = []
105 for k, v in src.named_parameters():
106 if 'ffn_ln' in k.lower():
107 src_ln.append((k, v))
108 dst_ln = []
109 for k, v in dst.named_parameters():
110 if 'ffn_ln' in k.lower():
111 dst_ln.append((k, v))
112 assert len(src_ln) == len(dst_ln)
113 for kvs, kvd in zip(src_ln, dst_ln):
114 assert kvd[1].data.shape == kvs[1].data.shape
115 kvd[1].data = kvs[1].data
116 assert (kvd[1].data == kvs[1].data).all()
117
118def copy_transformer_layer_wo_ln(src, dst, w2):
119 new_weight = torch.cat([src.attn.q_proj.weight.data, src.attn.k_proj.weight.data, src.attn.v_proj.weight.data], 0)

Callers 1

transform_weightFunction · 0.70

Calls 1

appendMethod · 0.80

Tested by

no test coverage detected