| 17 | |
| 18 | |
| 19 | def load_notes_encoder(weights, model): |
| 20 | model.token_embedder.weight = nn.Parameter(torch.FloatTensor(weights["token_embedder"]["embedding"])) |
| 21 | model.position_encoding.weight = nn.Parameter( |
| 22 | torch.FloatTensor(weights["Embed_0"]["embedding"]), requires_grad=False |
| 23 | ) |
| 24 | for lyr_num, lyr in enumerate(model.encoders): |
| 25 | ly_weight = weights[f"layers_{lyr_num}"] |
| 26 | lyr.layer[0].layer_norm.weight = nn.Parameter( |
| 27 | torch.FloatTensor(ly_weight["pre_attention_layer_norm"]["scale"]) |
| 28 | ) |
| 29 | |
| 30 | attention_weights = ly_weight["attention"] |
| 31 | lyr.layer[0].SelfAttention.q.weight = nn.Parameter(torch.FloatTensor(attention_weights["query"]["kernel"].T)) |
| 32 | lyr.layer[0].SelfAttention.k.weight = nn.Parameter(torch.FloatTensor(attention_weights["key"]["kernel"].T)) |
| 33 | lyr.layer[0].SelfAttention.v.weight = nn.Parameter(torch.FloatTensor(attention_weights["value"]["kernel"].T)) |
| 34 | lyr.layer[0].SelfAttention.o.weight = nn.Parameter(torch.FloatTensor(attention_weights["out"]["kernel"].T)) |
| 35 | |
| 36 | lyr.layer[1].layer_norm.weight = nn.Parameter(torch.FloatTensor(ly_weight["pre_mlp_layer_norm"]["scale"])) |
| 37 | |
| 38 | lyr.layer[1].DenseReluDense.wi_0.weight = nn.Parameter(torch.FloatTensor(ly_weight["mlp"]["wi_0"]["kernel"].T)) |
| 39 | lyr.layer[1].DenseReluDense.wi_1.weight = nn.Parameter(torch.FloatTensor(ly_weight["mlp"]["wi_1"]["kernel"].T)) |
| 40 | lyr.layer[1].DenseReluDense.wo.weight = nn.Parameter(torch.FloatTensor(ly_weight["mlp"]["wo"]["kernel"].T)) |
| 41 | |
| 42 | model.layer_norm.weight = nn.Parameter(torch.FloatTensor(weights["encoder_norm"]["scale"])) |
| 43 | return model |
| 44 | |
| 45 | |
| 46 | def load_continuous_encoder(weights, model): |