added fix

This commit is contained in:
Gang Wen
2026-04-06 22:01:39 +02:00
parent d55bd6acbb
commit e1dbbca003
2 changed files with 23 additions and 11 deletions
+7 -6
View File
@@ -249,6 +249,7 @@ class Florence2LanguageConfig(PretrainedConfig):
self.use_cache = use_cache
self.num_hidden_layers = encoder_layers
self.scale_embedding = scale_embedding # scale factor will be sqrt(d_model) if True
self.forced_bos_token_id = bos_token_id
super().__init__(
num_labels=num_labels,
@@ -262,12 +263,12 @@ class Florence2LanguageConfig(PretrainedConfig):
)
# ensure backward compatibility for BART CNN models
if self.forced_bos_token_id is None and kwargs.get("force_bos_token_to_be_generated", False):
self.forced_bos_token_id = self.bos_token_id
warnings.warn(
f"Please make sure the config includes `forced_bos_token_id={self.bos_token_id}` in future versions. "
"The config can simply be saved and uploaded again to be fixed."
)
# if self.forced_bos_token_id is None and kwargs.get("force_bos_token_to_be_generated", False):
# self.forced_bos_token_id = self.bos_token_id
# warnings.warn(
# f"Please make sure the config includes `forced_bos_token_id={self.bos_token_id}` in future versions. "
# "The config can simply be saved and uploaded again to be fixed."
# )
class Florence2Config(PretrainedConfig):
r"""
+15 -4
View File
@@ -61,6 +61,8 @@ from transformers.modeling_outputs import (
Seq2SeqLMOutput,
Seq2SeqModelOutput,
)
import transformers
from packaging import version
if is_flash_attn_2_available():
@@ -1934,7 +1936,8 @@ class Florence2Decoder(Florence2LanguagePreTrainedModel):
class Florence2LanguageModel(Florence2LanguagePreTrainedModel):
_tied_weights_keys = ["encoder.embed_tokens.weight", "decoder.embed_tokens.weight"]
if not version.parse(transformers.__version__) >= version.parse('5.0.0'):
_tied_weights_keys = ["encoder.embed_tokens.weight", "decoder.embed_tokens.weight"]
def __init__(self, config: Florence2LanguageConfig):
super().__init__(config)
@@ -2057,7 +2060,8 @@ class Florence2LanguageModel(Florence2LanguagePreTrainedModel):
class Florence2LanguageForConditionalGeneration(Florence2LanguagePreTrainedModel, GenerationMixin):
base_model_prefix = "model"
_tied_weights_keys = ["encoder.embed_tokens.weight", "decoder.embed_tokens.weight", "lm_head.weight"]
if not version.parse(transformers.__version__) >= version.parse('5.0.0'):
_tied_weights_keys = ["encoder.embed_tokens.weight", "decoder.embed_tokens.weight", "lm_head.weight"]
_keys_to_ignore_on_load_missing = ["final_logits_bias"]
def __init__(self, config: Florence2LanguageConfig):
@@ -2067,7 +2071,8 @@ class Florence2LanguageForConditionalGeneration(Florence2LanguagePreTrainedModel
self.lm_head = nn.Linear(config.d_model, self.model.shared.num_embeddings, bias=False)
# Initialize weights and apply final processing
self.post_init()
if not version.parse(transformers.__version__) >= version.parse('5.0.0'):
self.post_init()
def _tie_weights(self):
if self.config.tie_word_embeddings:
@@ -2075,6 +2080,11 @@ class Florence2LanguageForConditionalGeneration(Florence2LanguagePreTrainedModel
self._tie_or_clone_weights(self.model.decoder.embed_tokens, self.model.shared)
self._tie_or_clone_weights(self.lm_head, self.model.shared)
def tie_weights(self):
self.model.encoder.embed_tokens.weight = self.model.shared.weight
self.model.decoder.embed_tokens.weight = self.model.shared.weight
self.lm_head.weight = self.model.shared.weight
def get_encoder(self):
return self.model.get_encoder()
@@ -2536,7 +2546,8 @@ class Florence2VisionModelWithProjection(Florence2PreTrainedModel):
FLORENCE2_START_DOCSTRING,
)
class Florence2ForConditionalGeneration(Florence2PreTrainedModel, GenerationMixin):
_tied_weights_keys = ["language_model.encoder.embed_tokens.weight", "language_model.decoder.embed_tokens.weight", "language_model.lm_head.weight"]
if not version.parse(transformers.__version__) >= version.parse('5.0.0'):
_tied_weights_keys = ["language_model.encoder.embed_tokens.weight", "language_model.decoder.embed_tokens.weight", "language_model.lm_head.weight"]
def __init__(self, config: Florence2Config):
super().__init__(config)
assert config.vision_config.model_type == 'davit', 'only DaViT is supported for now'