MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / update_config

Function update_config

convert_to_hf.py:22–79  ·  view source on GitHub ↗
(
    source_config: dict,
    bos_token_id: int,
    eos_token_id: int,
    cls_token_id: int,
    pad_token_id: int,
    sep_token_id: int,
    max_length: int,
    torch_dtype: TorchDtype,
)

Source from the content-addressed store, hash-verified

20
21
22def 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

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected