TheoDB commited on
Commit
f3db5c9
·
verified ·
1 Parent(s): 0afa2d1

Add transformers v5 support on main; persist full sliding_window; tidy auto_map

Browse files
Files changed (4) hide show
  1. README.md +11 -1
  2. config.json +4 -5
  3. configuration_bidirlm.py +20 -7
  4. modeling_bidirlm.py +33 -19
README.md CHANGED
@@ -119,11 +119,21 @@ mlm = AutoModelForMaskedLM.from_pretrained("BidirLM/BidirLM-270M-Base", trust_re
119
  ## Requirements
120
 
121
  ```
122
- transformers>=4.57.6,<5.0.0
123
  ```
124
 
125
  This model requires `trust_remote_code=True`.
126
 
 
 
 
 
 
 
 
 
 
 
127
  ## Citation
128
 
129
  ```bibtex
 
119
  ## Requirements
120
 
121
  ```
122
+ transformers>=5.0
123
  ```
124
 
125
  This model requires `trust_remote_code=True`.
126
 
127
+ > **Note:** This model was trained with `transformers==4.57.6` (transformers 4.x). The version on `main` was patched to work with `transformers>=5.0`. For the original (pre-patch) version, which is compatible with `transformers>=4.57.6,<5.0.0`, use the `transformers-v4` branch:
128
+ > ```python
129
+ > from transformers import AutoModel
130
+ > model = AutoModel.from_pretrained(
131
+ > "BidirLM/BidirLM-270M-Base",
132
+ > trust_remote_code=True,
133
+ > revision="transformers-v4",
134
+ > )
135
+ > ```
136
+
137
  ## Citation
138
 
139
  ```bibtex
config.json CHANGED
@@ -7,10 +7,9 @@
7
  "attention_dropout": 0.0,
8
  "attn_logit_softcapping": null,
9
  "auto_map": {
10
- "AutoConfig": "configuration_bidirlm.BidirLMTextConfig",
11
- "AutoModel": "modeling_bidirlm.BidirLMTextModel",
12
  "AutoModelForMaskedLM": "modeling_bidirlm.BidirLMForMaskedLM",
13
- "AutoModelForPreTraining": "modeling_bidirlm.BidirLMPreTrainedModel",
14
  "AutoModelForSequenceClassification": "modeling_bidirlm.BidirLMForSequenceClassification",
15
  "AutoModelForTokenClassification": "modeling_bidirlm.BidirLMForTokenClassification"
16
  },
@@ -56,8 +55,8 @@
56
  "rope_scaling": null,
57
  "rope_theta": 1000000.0,
58
  "sliding_window": 512,
59
- "transformers_version": "4.57.3",
60
  "use_bidirectional_attention": true,
61
  "use_cache": true,
62
  "vocab_size": 262144
63
- }
 
7
  "attention_dropout": 0.0,
8
  "attn_logit_softcapping": null,
9
  "auto_map": {
10
+ "AutoConfig": "configuration_bidirlm.BidirLMConfig",
11
+ "AutoModel": "modeling_bidirlm.BidirLMModel",
12
  "AutoModelForMaskedLM": "modeling_bidirlm.BidirLMForMaskedLM",
 
13
  "AutoModelForSequenceClassification": "modeling_bidirlm.BidirLMForSequenceClassification",
14
  "AutoModelForTokenClassification": "modeling_bidirlm.BidirLMForTokenClassification"
15
  },
 
55
  "rope_scaling": null,
56
  "rope_theta": 1000000.0,
57
  "sliding_window": 512,
58
+ "transformers_version": "5.9.0",
59
  "use_bidirectional_attention": true,
60
  "use_cache": true,
61
  "vocab_size": 262144
62
+ }
configuration_bidirlm.py CHANGED
@@ -23,13 +23,14 @@ from typing import Any, Optional, Union
23
 
24
  import transformers
25
  _v = transformers.__version__
26
- if _v < "4.57.6" or _v >= "5.0.0":
27
  raise ImportError(
28
- f"BidirLM requires transformers>=4.57.6,<5.0.0 (found {_v}). "
29
- f"Install a compatible version: pip install 'transformers>=4.57.6,<5.0.0'"
 
30
  )
31
 
32
- from transformers.configuration_utils import PretrainedConfig, layer_type_validation
33
  from transformers.modeling_rope_utils import rope_config_validation
34
  from transformers.utils import logging
35
  from transformers.models.siglip import SiglipVisionConfig
@@ -174,6 +175,20 @@ class BidirLMConfig(PretrainedConfig):
174
  "norm": (["hidden_states"], ["hidden_states"]),
175
  }
176
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
177
  def __init__(
178
  self,
179
  vocab_size=262_208,
@@ -235,8 +250,6 @@ class BidirLMConfig(PretrainedConfig):
235
  self.layer_types = layer_types
236
  self.use_bidirectional_attention = use_bidirectional_attention
237
  self.classifier_pooling = classifier_pooling
238
- if use_bidirectional_attention:
239
- self.sliding_window = self.sliding_window // 2
240
 
241
  self.rope_local_base_freq = rope_local_base_freq
242
  self.rope_scaling = rope_scaling
@@ -250,7 +263,7 @@ class BidirLMConfig(PretrainedConfig):
250
  "sliding_attention" if bool((i + 1) % self._sliding_window_pattern) else "full_attention"
251
  for i in range(self.num_hidden_layers)
252
  ]
253
- layer_type_validation(self.layer_types, self.num_hidden_layers)
254
 
255
 
256
  class Gemma3Config(PretrainedConfig):
 
23
 
24
  import transformers
25
  _v = transformers.__version__
26
+ if _v < "5.0.0":
27
  raise ImportError(
28
+ f"BidirLM requires transformers>=5.0.0 on this branch (found {_v}). "
29
+ f"Install a compatible version: pip install 'transformers>=5.0.0'. "
30
+ f"For transformers 4.x, use the `transformers-v4` branch instead."
31
  )
32
 
33
+ from transformers.configuration_utils import PretrainedConfig
34
  from transformers.modeling_rope_utils import rope_config_validation
35
  from transformers.utils import logging
36
  from transformers.models.siglip import SiglipVisionConfig
 
175
  "norm": (["hidden_states"], ["hidden_states"]),
176
  }
177
 
178
+ @property
179
+ def effective_sliding_window(self):
180
+ # Per-side sliding window used at attention time. For bidirectional
181
+ # models the configured (total) window is split symmetrically, so
182
+ # each side is half. Derived at runtime so the full `sliding_window`
183
+ # is what gets persisted (no halving on every save/load).
184
+ sw = getattr(self, "sliding_window", None)
185
+ if sw is None:
186
+ return None
187
+ if getattr(self, "use_bidirectional_attention", False):
188
+ return sw // 2
189
+ return sw
190
+
191
+
192
  def __init__(
193
  self,
194
  vocab_size=262_208,
 
250
  self.layer_types = layer_types
251
  self.use_bidirectional_attention = use_bidirectional_attention
252
  self.classifier_pooling = classifier_pooling
 
 
253
 
254
  self.rope_local_base_freq = rope_local_base_freq
255
  self.rope_scaling = rope_scaling
 
263
  "sliding_attention" if bool((i + 1) % self._sliding_window_pattern) else "full_attention"
264
  for i in range(self.num_hidden_layers)
265
  ]
266
+ self.validate_layer_type()
267
 
268
 
269
  class Gemma3Config(PretrainedConfig):
modeling_bidirlm.py CHANGED
@@ -3,10 +3,11 @@ from typing import Optional
3
 
4
  import transformers
5
  _v = transformers.__version__
6
- if _v < "4.57.6" or _v >= "5.0.0":
7
  raise ImportError(
8
- f"BidirLM requires transformers>=4.57.6,<5.0.0 (found {_v}). "
9
- f"Install a compatible version: pip install 'transformers>=4.57.6,<5.0.0'"
 
10
  )
11
 
12
  import torch
@@ -126,7 +127,7 @@ class Gemma3Attention(nn.Module):
126
  bias=config.attention_bias,
127
  )
128
  self.attn_logit_softcapping = self.config.attn_logit_softcapping
129
- self.sliding_window = config.sliding_window if self.is_sliding else None
130
 
131
  self.q_norm = Gemma3RMSNorm(dim=config.head_dim, eps=config.rms_norm_eps)
132
  self.k_norm = Gemma3RMSNorm(dim=config.head_dim, eps=config.rms_norm_eps)
@@ -334,13 +335,12 @@ class BidirLMPreTrainedModel(PreTrainedModel):
334
 
335
  def _init_weights(self, module):
336
  super()._init_weights(module)
337
- # if isinstance(module, Gemma3MultiModalProjector):
338
- # module.mm_input_projection_weight.data.zero_()
339
- # # We initialize with 0s to be 1 centered as the RMSNorm here does (1 + weight)
340
- # elif "RMSNorm" in module.__class__.__name__:
341
- # module.weight.data.zero_()
342
  if "RMSNorm" in module.__class__.__name__:
343
  module.weight.data.zero_()
 
 
 
 
344
 
345
 
346
  class Gemma3TextScaledWordEmbedding(nn.Embedding):
@@ -356,6 +356,10 @@ class Gemma3TextScaledWordEmbedding(nn.Embedding):
356
  embed_scale: float = 1.0,
357
  ):
358
  super().__init__(num_embeddings, embedding_dim, padding_idx)
 
 
 
 
359
  self.register_buffer("embed_scale", torch.tensor(embed_scale), persistent=False)
360
 
361
  def forward(self, input_ids: torch.Tensor):
@@ -411,12 +415,21 @@ class Gemma3RotaryEmbedding(nn.Module):
411
  self.original_max_seq_len = config.max_position_embeddings
412
 
413
  self.config = config
414
- self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
 
415
 
416
- inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
417
  self.register_buffer("inv_freq", inv_freq, persistent=False)
418
  self.original_inv_freq = self.inv_freq
419
 
 
 
 
 
 
 
 
 
420
  @torch.no_grad()
421
  @dynamic_rope_update
422
  def forward(self, x, position_ids):
@@ -541,7 +554,7 @@ class BidirLMModel(BidirLMPreTrainedModel):
541
  else self.config.output_hidden_states
542
  )
543
  return_dict = (
544
- return_dict if return_dict is not None else self.config.use_return_dict
545
  )
546
  all_hidden_states = () if output_hidden_states else None
547
  all_self_attns = () if output_attentions else None
@@ -573,10 +586,11 @@ class BidirLMModel(BidirLMPreTrainedModel):
573
  position_embeddings_global = self.rotary_emb(hidden_states, position_ids)
574
  position_embeddings_local = self.rotary_emb_local(hidden_states, position_ids)
575
 
 
576
  window_size = (
577
  (
578
- self.config.sliding_window,
579
- self.config.sliding_window if self.config.use_bidirectional_attention else 0
580
  )
581
  if self.config.sliding_window is not None
582
  else None
@@ -645,7 +659,7 @@ class BidirLMModel(BidirLMPreTrainedModel):
645
 
646
 
647
  class BidirLMForMaskedLM(BidirLMPreTrainedModel):
648
- _tied_weights_keys = ["lm_head.weight"]
649
  config: BidirLMConfig
650
 
651
  def __init__(self, config):
@@ -670,7 +684,7 @@ class BidirLMForMaskedLM(BidirLMPreTrainedModel):
670
  **kwargs,
671
  ) -> tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
672
  return_dict = (
673
- return_dict if return_dict is not None else self.config.use_return_dict
674
  )
675
  encoder_output = self.model(
676
  input_ids=input_ids,
@@ -726,7 +740,7 @@ class BidirLMForSequenceClassification(BidirLMPreTrainedModel):
726
  **kwargs,
727
  ) -> tuple[torch.Tensor] | SequenceClassifierOutput:
728
  return_dict = (
729
- return_dict if return_dict is not None else self.config.use_return_dict
730
  )
731
 
732
  encoder_output = self.model(
@@ -823,7 +837,7 @@ class BidirLMForTokenClassification(BidirLMPreTrainedModel):
823
  return_dict: Optional[bool] = None,
824
  ) -> tuple[torch.Tensor] | TokenClassifierOutput:
825
  return_dict = (
826
- return_dict if return_dict is not None else self.config.use_return_dict
827
  )
828
 
829
  outputs = self.model(
@@ -976,7 +990,7 @@ class BidirLMForTokenClassification(BidirLMPreTrainedModel):
976
  # output_hidden_states = (
977
  # output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
978
  # )
979
- # return_dict = return_dict if return_dict is not None else self.config.use_return_dict
980
 
981
  # # Replace image id with PAD if the image token if OOV, to avoid index-errors
982
  # if input_ids is not None and self.config.image_token_id >= self.vocab_size:
 
3
 
4
  import transformers
5
  _v = transformers.__version__
6
+ if _v < "5.0.0":
7
  raise ImportError(
8
+ f"BidirLM requires transformers>=5.0.0 on this branch (found {_v}). "
9
+ f"Install a compatible version: pip install 'transformers>=5.0.0'. "
10
+ f"For transformers 4.x, use the `transformers-v4` branch instead."
11
  )
12
 
13
  import torch
 
127
  bias=config.attention_bias,
128
  )
129
  self.attn_logit_softcapping = self.config.attn_logit_softcapping
130
+ self.sliding_window = config.effective_sliding_window if self.is_sliding else None
131
 
132
  self.q_norm = Gemma3RMSNorm(dim=config.head_dim, eps=config.rms_norm_eps)
133
  self.k_norm = Gemma3RMSNorm(dim=config.head_dim, eps=config.rms_norm_eps)
 
335
 
336
  def _init_weights(self, module):
337
  super()._init_weights(module)
 
 
 
 
 
338
  if "RMSNorm" in module.__class__.__name__:
339
  module.weight.data.zero_()
340
+ elif isinstance(module, Gemma3TextScaledWordEmbedding):
341
+ # transformers 5.x resets non-persistent buffers in this hook;
342
+ # restore embed_scale to the value chosen at construction.
343
+ torch.nn.init.constant_(module.embed_scale, module.scalar_embed_scale)
344
 
345
 
346
  class Gemma3TextScaledWordEmbedding(nn.Embedding):
 
356
  embed_scale: float = 1.0,
357
  ):
358
  super().__init__(num_embeddings, embedding_dim, padding_idx)
359
+ # transformers 5.x calls _init_weights on every module post-load and
360
+ # resets non-persistent buffers to zero; keep the scalar so the
361
+ # PreTrainedModel._init_weights below can re-initialize embed_scale.
362
+ self.scalar_embed_scale = embed_scale
363
  self.register_buffer("embed_scale", torch.tensor(embed_scale), persistent=False)
364
 
365
  def forward(self, input_ids: torch.Tensor):
 
415
  self.original_max_seq_len = config.max_position_embeddings
416
 
417
  self.config = config
418
+ # transformers 5.x removed 'default' from ROPE_INIT_FUNCTIONS
419
+ rope_init_fn = self.compute_default_rope_parameters if self.rope_type == "default" else ROPE_INIT_FUNCTIONS[self.rope_type]
420
 
421
+ inv_freq, self.attention_scaling = rope_init_fn(self.config, device)
422
  self.register_buffer("inv_freq", inv_freq, persistent=False)
423
  self.original_inv_freq = self.inv_freq
424
 
425
+ @staticmethod
426
+ def compute_default_rope_parameters(config, device=None, **kwargs):
427
+ base = config.rope_theta
428
+ head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
429
+ dim = int(head_dim * getattr(config, "partial_rotary_factor", 1.0))
430
+ inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim))
431
+ return inv_freq, 1.0
432
+
433
  @torch.no_grad()
434
  @dynamic_rope_update
435
  def forward(self, x, position_ids):
 
554
  else self.config.output_hidden_states
555
  )
556
  return_dict = (
557
+ return_dict if return_dict is not None else True
558
  )
559
  all_hidden_states = () if output_hidden_states else None
560
  all_self_attns = () if output_attentions else None
 
586
  position_embeddings_global = self.rotary_emb(hidden_states, position_ids)
587
  position_embeddings_local = self.rotary_emb_local(hidden_states, position_ids)
588
 
589
+ _eff_sw = self.config.effective_sliding_window
590
  window_size = (
591
  (
592
+ _eff_sw,
593
+ _eff_sw if self.config.use_bidirectional_attention else 0
594
  )
595
  if self.config.sliding_window is not None
596
  else None
 
659
 
660
 
661
  class BidirLMForMaskedLM(BidirLMPreTrainedModel):
662
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
663
  config: BidirLMConfig
664
 
665
  def __init__(self, config):
 
684
  **kwargs,
685
  ) -> tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
686
  return_dict = (
687
+ return_dict if return_dict is not None else True
688
  )
689
  encoder_output = self.model(
690
  input_ids=input_ids,
 
740
  **kwargs,
741
  ) -> tuple[torch.Tensor] | SequenceClassifierOutput:
742
  return_dict = (
743
+ return_dict if return_dict is not None else True
744
  )
745
 
746
  encoder_output = self.model(
 
837
  return_dict: Optional[bool] = None,
838
  ) -> tuple[torch.Tensor] | TokenClassifierOutput:
839
  return_dict = (
840
+ return_dict if return_dict is not None else True
841
  )
842
 
843
  outputs = self.model(
 
990
  # output_hidden_states = (
991
  # output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
992
  # )
993
+ # return_dict = return_dict if return_dict is not None else True
994
 
995
  # # Replace image id with PAD if the image token if OOV, to avoid index-errors
996
  # if input_ids is not None and self.config.image_token_id >= self.vocab_size: