| 20 | |
| 21 | |
| 22 | def update_config( |
| 23 | source_config: dict, |
| 24 | bos_token_id: int, |
| 25 | eos_token_id: int, |
| 26 | cls_token_id: int, |
| 27 | pad_token_id: int, |
| 28 | sep_token_id: int, |
| 29 | max_length: int, |
| 30 | torch_dtype: TorchDtype, |
| 31 | ) -> dict: |
| 32 | target_config = { |
| 33 | # "_name_or_path": "ModernBERT-base", |
| 34 | "architectures": ["ModernBertForMaskedLM"], |
| 35 | "attention_bias": source_config["attn_out_bias"], |
| 36 | "attention_dropout": source_config["attention_probs_dropout_prob"], |
| 37 | "bos_token_id": bos_token_id, |
| 38 | "classifier_activation": source_config.get("head_class_act", source_config["hidden_act"]), |
| 39 | "classifier_bias": source_config["head_class_bias"], |
| 40 | "classifier_dropout": source_config["head_class_dropout"], |
| 41 | "classifier_pooling": "mean", |
| 42 | "cls_token_id": cls_token_id, |
| 43 | "decoder_bias": source_config["decoder_bias"], |
| 44 | "deterministic_flash_attn": source_config["deterministic_fa2"], |
| 45 | "embedding_dropout": source_config["embed_dropout_prob"], |
| 46 | "eos_token_id": eos_token_id, |
| 47 | "global_attn_every_n_layers": source_config["global_attn_every_n_layers"], |
| 48 | "global_rope_theta": source_config["rotary_emb_base"], |
| 49 | "gradient_checkpointing": source_config["gradient_checkpointing"], |
| 50 | "hidden_activation": source_config["hidden_act"], |
| 51 | "hidden_size": source_config["hidden_size"], |
| 52 | "initializer_cutoff_factor": source_config["init_cutoff_factor"], |
| 53 | "initializer_range": source_config["initializer_range"], |
| 54 | "intermediate_size": source_config["intermediate_size"], |
| 55 | "layer_norm_eps": source_config["norm_kwargs"]["eps"], |
| 56 | "local_attention": source_config["sliding_window"], |
| 57 | "local_rope_theta": source_config["local_attn_rotary_emb_base"] |
| 58 | if ( |
| 59 | source_config["local_attn_rotary_emb_base"] |
| 60 | and source_config["local_attn_rotary_emb_base"] != -1 |
| 61 | ) |
| 62 | else source_config["rotary_emb_base"], |
| 63 | "max_position_embeddings": max_length, # Override with first config value |
| 64 | "mlp_bias": source_config["mlp_in_bias"], |
| 65 | "mlp_dropout": source_config["mlp_dropout_prob"], |
| 66 | "model_type": "modernbert", |
| 67 | "norm_bias": source_config["norm_kwargs"]["bias"], |
| 68 | "norm_eps": source_config["norm_kwargs"]["eps"], |
| 69 | "num_attention_heads": source_config["num_attention_heads"], |
| 70 | "num_hidden_layers": source_config["num_hidden_layers"], |
| 71 | "pad_token_id": pad_token_id, |
| 72 | "position_embedding_type": source_config["position_embedding_type"], |
| 73 | "sep_token_id": sep_token_id, |
| 74 | "tie_word_embeddings": source_config.get("tie_word_embeddings", True), |
| 75 | "torch_dtype": torch_dtype.value, |
| 76 | "transformers_version": "4.48.0", |
| 77 | "vocab_size": source_config["vocab_size"], |
| 78 | } |
| 79 | return target_config |