(src, dst)
| 52 | dst.data = src.data |
| 53 | |
| 54 | def copy_layer_norm(src, dst): |
| 55 | src_ln = [] |
| 56 | for k, v in src.named_parameters(): |
| 57 | if 'norm' in k.lower() and type(v) is not torch.nn.Identity(): |
| 58 | src_ln.append((k, v)) |
| 59 | dst_ln = [] |
| 60 | for k, v in dst.named_parameters(): |
| 61 | if 'layernorm' in k.lower(): |
| 62 | dst_ln.append((k, v)) |
| 63 | assert len(src_ln) == len(dst_ln) |
| 64 | for kvs, kvd in zip(src_ln, dst_ln): |
| 65 | assert kvd[1].data.shape == kvs[1].data.shape |
| 66 | kvd[1].data = kvs[1].data |
| 67 | assert (kvd[1].data == kvs[1].data).all() |
| 68 | |
| 69 | def copy_transformer_layer_wo_ln(src, dst): |
| 70 | new_weight = src.attn.qkv.weight.data |
no test coverage detected