(src_model, swiss_model)
| 66 | copy_layer_param(src.mlp.fc2, dst.mlp.dense_4h_to_h) |
| 67 | |
| 68 | def transform_weight(src_model, swiss_model): |
| 69 | # transform layernorm |
| 70 | copy_layer_norm(src_model, swiss_model) |
| 71 | # transform encoder |
| 72 | copy_from_param(src_model.cls_token.data[0], swiss_model.encoder.transformer.word_embeddings.weight) |
| 73 | copy_from_param(src_model.pos_embed.data[0], swiss_model.encoder.transformer.position_embeddings.weight) |
| 74 | for src_l, dst_l in zip(src_model.blocks, swiss_model.encoder.transformer.layers): |
| 75 | copy_transformer_layer_wo_ln(src_l, dst_l) |
| 76 | copy_layer_param(src_model.patch_embed.proj, swiss_model.encoder.mixins.patch_embedding.proj) |
| 77 | # transform decoder |
| 78 | copy_from_param(src_model.mask_token.data[0], swiss_model.decoder.transformer.word_embeddings.weight) |
| 79 | copy_from_param(src_model.decoder_pos_embed.data[0], swiss_model.decoder.transformer.position_embeddings.weight) |
| 80 | for src_l, dst_l in zip(src_model.decoder_blocks, swiss_model.decoder.transformer.layers): |
| 81 | copy_transformer_layer_wo_ln(src_l, dst_l) |
| 82 | copy_layer_param(src_model.decoder_embed, swiss_model.decoder.mixins.mask_forward.decoder_embed) |
| 83 | copy_layer_param(src_model.decoder_pred, swiss_model.decoder.mixins.mask_forward.decoder_pred) |
| 84 | |
| 85 | |
| 86 | vit.eval() |
no test coverage detected