added fix
This commit is contained in:
@@ -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"""
|
||||
|
||||
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user