(src, dst)
| 88 | assert (dst_dic[k].data == src_dic[k].data).all() |
| 89 | |
| 90 | def 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 | |
| 118 | def 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) |
no test coverage detected