(src, dst)
| 40 | dst.data = src.data |
| 41 | |
| 42 | def copy_layer_norm(src, dst): |
| 43 | src_ln = [] |
| 44 | for k, v in src.named_parameters(): |
| 45 | if 'norm' in k.lower() and type(v) is not torch.nn.Identity(): |
| 46 | src_ln.append((k, v)) |
| 47 | dst_ln = [] |
| 48 | for k, v in dst.named_parameters(): |
| 49 | if 'layernorm' in k.lower(): |
| 50 | dst_ln.append((k, v)) |
| 51 | assert len(src_ln) == len(dst_ln) |
| 52 | for kvs, kvd in zip(src_ln, dst_ln): |
| 53 | assert kvd[1].data.shape == kvs[1].data.shape |
| 54 | kvd[1].data = kvs[1].data |
| 55 | assert (kvd[1].data == kvs[1].data).all() |
| 56 | |
| 57 | def copy_transformer_layer_wo_ln(src, dst): |
| 58 | new_weight = src.attn.qkv.weight.data |
no test coverage detected