(self, spec, module, rotary_dim, num_heads)
| 738 | config.unk_token = tokenizer.unk_token |
| 739 | |
| 740 | def set_decoder(self, spec, module, rotary_dim, num_heads): |
| 741 | spec.scale_embeddings = False |
| 742 | self.set_embeddings(spec.embeddings, module.wte) |
| 743 | self.set_layer_norm(spec.layer_norm, module.ln_f) |
| 744 | |
| 745 | for layer_spec, layer in zip(spec.layer, module.h): |
| 746 | self.set_layer_norm(layer_spec.shared_layer_norm, layer.ln_1) |
| 747 | |
| 748 | qw = layer.attn.q_proj.weight |
| 749 | kw = layer.attn.k_proj.weight |
| 750 | vw = layer.attn.v_proj.weight |
| 751 | |
| 752 | qw = utils.permute_for_sliced_rotary(qw, num_heads, rotary_dim) |
| 753 | kw = utils.permute_for_sliced_rotary(kw, num_heads, rotary_dim) |
| 754 | |
| 755 | layer_spec.self_attention.linear[0].weight = torch.cat((qw, kw, vw)) |
| 756 | self.set_linear(layer_spec.self_attention.linear[1], layer.attn.out_proj) |
| 757 | |
| 758 | self.set_linear(layer_spec.ffn.linear_0, layer.mlp.fc_in) |
| 759 | self.set_linear(layer_spec.ffn.linear_1, layer.mlp.fc_out) |
| 760 | |
| 761 | |
| 762 | @register_loader("CodeGenConfig") |
no test coverage detected