Skip to content
Merged
Show file tree
Hide file tree
Changes from 10 commits
Commits
Show all changes
74 commits
Select commit Hold shift + click to select a range
028e502
first try
zucchini-nlp Jul 25, 2024
589d18a
codestyle
zucchini-nlp Jul 25, 2024
b33982f
idefics2 is happy
zucchini-nlp Jul 26, 2024
b0baa75
[run-slow] llava, llava_next, video_llava, vipllava, llava_next_video…
zucchini-nlp Jul 26, 2024
9f19211
fix-copies
zucchini-nlp Jul 26, 2024
19e0f3f
[run-slow] llava, llava_next, video_llava, vipllava, llava_next_video…
zucchini-nlp Jul 26, 2024
56a6f81
blip-2 needs to init vision from config
zucchini-nlp Jul 26, 2024
bbff1ac
when was this removed O_o
zucchini-nlp Jul 26, 2024
8485df9
minor fix
zucchini-nlp Jul 26, 2024
ba7ee7f
tests
zucchini-nlp Jul 26, 2024
d793d04
this way?
zucchini-nlp Jul 29, 2024
58aff27
tests
zucchini-nlp Jul 29, 2024
4e4bd27
model-agnostic code
zucchini-nlp Jul 31, 2024
52e77b3
codestyle
zucchini-nlp Jul 31, 2024
0feec1e
add tests for idefics
zucchini-nlp Jul 31, 2024
7b61096
modify general test for VLMs
zucchini-nlp Jul 31, 2024
9d15024
no generation test for vlm yet!
zucchini-nlp Jul 31, 2024
4e20af6
no generation test here also
zucchini-nlp Jul 31, 2024
1c6435d
wanr in VIT-SDPA if output attn
zucchini-nlp Jul 31, 2024
ea71d89
add more tests
zucchini-nlp Jul 31, 2024
17f9e69
user can pass dict as attn impl
zucchini-nlp Jul 31, 2024
140f222
repo consistency
zucchini-nlp Jul 31, 2024
3648537
update
zucchini-nlp Aug 1, 2024
ca95fee
muicgen
zucchini-nlp Aug 1, 2024
79cae6d
no prints
zucchini-nlp Aug 1, 2024
378274b
forgot speech enc-dec and clip
zucchini-nlp Aug 1, 2024
a772ff5
how many composite models we have?
zucchini-nlp Aug 1, 2024
3e6787c
musicgen meelody is same as mudicgen
zucchini-nlp Aug 1, 2024
00b2065
+siglip
zucchini-nlp Aug 1, 2024
5cdfbfb
fix tests + add some more
zucchini-nlp Aug 2, 2024
723f27d
remove idefics custom overriden code
zucchini-nlp Aug 2, 2024
198c60c
make idefics2 automappable
zucchini-nlp Aug 7, 2024
2713616
nits
zucchini-nlp Aug 7, 2024
3aef763
skip tests
zucchini-nlp Aug 7, 2024
6c31934
doctests
zucchini-nlp Aug 7, 2024
2f90176
Update src/transformers/models/idefics2/configuration_idefics2.py
zucchini-nlp Aug 8, 2024
3dfb48c
Update tests/models/clip/test_modeling_clip.py
zucchini-nlp Aug 8, 2024
505ed3f
Update tests/models/idefics2/test_modeling_idefics2.py
zucchini-nlp Aug 8, 2024
4c9f894
Update tests/models/idefics2/test_modeling_idefics2.py
zucchini-nlp Aug 8, 2024
d7d54f8
Update src/transformers/configuration_utils.py
zucchini-nlp Aug 8, 2024
b6e9951
major update, no need for automap
zucchini-nlp Aug 8, 2024
d1a291c
clean up
zucchini-nlp Aug 9, 2024
57119b1
add FA2 test
zucchini-nlp Aug 9, 2024
cfb9198
more tests
zucchini-nlp Aug 9, 2024
720694f
merge main
zucchini-nlp Aug 9, 2024
a2e9062
style
zucchini-nlp Aug 9, 2024
e72c03d
skip tests
zucchini-nlp Aug 9, 2024
60df872
why did these started failing now?
zucchini-nlp Aug 9, 2024
6d12897
no attributes for FA2 needed
zucchini-nlp Aug 9, 2024
957e64e
one tiny test
zucchini-nlp Aug 9, 2024
94e7578
address comment about FA2 false warning
zucchini-nlp Sep 18, 2024
2607d94
merge main
zucchini-nlp Oct 3, 2024
59aa480
style
zucchini-nlp Oct 3, 2024
1dc8bd1
add new models and resolve conflicts
zucchini-nlp Oct 3, 2024
04aba9f
fix copies
zucchini-nlp Oct 3, 2024
598b6f5
let it be this way for now, come back tomorrow to review
zucchini-nlp Oct 3, 2024
6d02e5c
some more fixes
zucchini-nlp Oct 4, 2024
a578fdc
update
zucchini-nlp Oct 4, 2024
c02a943
more updates
zucchini-nlp Oct 4, 2024
19de595
update
zucchini-nlp Oct 4, 2024
d691097
fix copies
zucchini-nlp Oct 4, 2024
1fa548f
Merge remote-tracking branch 'upstream/main' into vlm_sdpa_flag
zucchini-nlp Oct 4, 2024
39c032e
style and tests
zucchini-nlp Oct 4, 2024
74b211f
another big update
zucchini-nlp Oct 8, 2024
6dcf21a
fix tests
zucchini-nlp Oct 8, 2024
5a6ffe7
Merge remote-tracking branch 'upstream/main' into vlm_sdpa_flag
zucchini-nlp Oct 8, 2024
b93b79a
fix tests
zucchini-nlp Oct 9, 2024
702dacf
update
zucchini-nlp Oct 10, 2024
0616732
another update
zucchini-nlp Oct 10, 2024
28859ce
Merge branch 'main' into vlm_sdpa_flag
zucchini-nlp Oct 10, 2024
ebf15ef
fix tests
zucchini-nlp Oct 11, 2024
0941fee
merge main
zucchini-nlp Oct 21, 2024
24a9fc5
fix copies
zucchini-nlp Oct 21, 2024
4eed237
fix tests
zucchini-nlp Oct 21, 2024
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions src/transformers/modeling_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
from .dynamic_module_utils import custom_object_save
from .generation import GenerationConfig, GenerationMixin
from .integrations import PeftAdapterMixin, deepspeed_config, is_deepspeed_zero3_enabled
from .models.auto.modeling_auto import MODEL_MAPPING
from .pytorch_utils import ( # noqa: F401
Conv1D,
apply_chunking_to_forward,
Expand Down Expand Up @@ -1512,6 +1513,40 @@ def _autoset_attn_implementation(
# If a config is passed with a preset attn_implementation, we skip the automatic dispatch and use the user-provided config, with hard checks that the requested attention implementation is available.
requested_attn_implementation = config._attn_implementation_internal

# MultiModal LLM related hack: since they consist of two or might be more sub-models
# we have to check and dispatch SDPA to each sub-model, in case any of them support it.
# If one sub-model supports SDPA while other doesn't, an error will be raised following the
# typical SDPA-dispatch path.
# Same goes for the `_supports_cache_class` because we cannot know what flag LM has
# before knowing which class is that LM
Comment thread
zucchini-nlp marked this conversation as resolved.
Outdated
if hasattr(config, "text_config") and hasattr(config, "vision_config"):
text_model_cls = MODEL_MAPPING.get(type(config.text_config), None)
if text_model_cls is not None:
config.text_config._attn_implementation = config._attn_implementation
cls._supports_sdpa = text_model_cls._supports_sdpa
Comment thread
zucchini-nlp marked this conversation as resolved.
Outdated
cls._autoset_attn_implementation(
config.text_config,
use_flash_attention_2=use_flash_attention_2,
torch_dtype=torch_dtype,
device_map=device_map,
check_device_map=check_device_map,
)

vision_model_cls = MODEL_MAPPING.get(type(config.vision_config), None)
if vision_model_cls is not None:
config.vision_config._attn_implementation = config._attn_implementation
cls._supports_sdpa = vision_model_cls._supports_sdpa
cls._autoset_attn_implementation(
Comment thread
zucchini-nlp marked this conversation as resolved.
Outdated
config.vision_config,
use_flash_attention_2=use_flash_attention_2,
torch_dtype=torch_dtype,
device_map=device_map,
check_device_map=check_device_map,
)

if vision_model_cls is not None and text_model_cls is not None:
cls._supports_sdpa = vision_model_cls._supports_sdpa or text_model_cls._supports_sdpa
Comment thread
zucchini-nlp marked this conversation as resolved.
Outdated

if use_flash_attention_2:
logger.warning_once(
'The model was loaded with use_flash_attention_2=True, which is deprecated and may be removed in a future release. Please use `attn_implementation="flash_attention_2"` instead.'
Expand Down
12 changes: 8 additions & 4 deletions src/transformers/models/blip_2/modeling_blip_2.py
Original file line number Diff line number Diff line change
Expand Up @@ -1225,19 +1225,21 @@ class Blip2Model(Blip2PreTrainedModel):
def __init__(self, config: Blip2Config):
super().__init__(config)

self.vision_model = Blip2VisionModel(config.vision_config)
self.vision_model = Blip2VisionModel._from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)

self.query_tokens = nn.Parameter(torch.zeros(1, config.num_query_tokens, config.qformer_config.hidden_size))
self.qformer = Blip2QFormerModel(config.qformer_config)

self.language_projection = nn.Linear(config.qformer_config.hidden_size, config.text_config.hidden_size)
if config.use_decoder_only_language_model:
language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)
else:
language_model = AutoModelForSeq2SeqLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)

# Update _tied_weights_keys using the base model used.
Expand Down Expand Up @@ -1590,7 +1592,9 @@ class Blip2ForConditionalGeneration(Blip2PreTrainedModel):
def __init__(self, config: Blip2Config):
super().__init__(config)

self.vision_model = Blip2VisionModel(config.vision_config)
self.vision_model = Blip2VisionModel._from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)

self.query_tokens = nn.Parameter(torch.zeros(1, config.num_query_tokens, config.qformer_config.hidden_size))
self.qformer = Blip2QFormerModel(config.qformer_config)
Expand Down
17 changes: 15 additions & 2 deletions src/transformers/models/idefics2/modeling_idefics2.py
Original file line number Diff line number Diff line change
Expand Up @@ -921,7 +921,9 @@ def __init__(self, config, layer_idx: int):

self.input_latents_norm = Idefics2RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self.input_context_norm = Idefics2RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self.self_attn = IDEFICS2_PERCEIVER_ATTENTION_CLASSES[config._attn_implementation](config, layer_idx=layer_idx)
self.self_attn = IDEFICS2_PERCEIVER_ATTENTION_CLASSES[config.perceiver_config._attn_implementation](
config, layer_idx=layer_idx
)
self.post_attention_layernorm = Idefics2RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self.mlp = Idefics2MLP(
hidden_size=config.text_config.hidden_size,
Expand Down Expand Up @@ -1131,7 +1133,18 @@ def _autoset_attn_implementation(
check_device_map=check_device_map,
**kwargs,
)
config.vision_config._attn_implementation = config._attn_implementation
# autoset-attn calls recursively all sub-configs (text-config, vision-config)
# and sets attn implementation if the config can be mapped bu auto-model
# Idefics2 vision config can't be mapped automcatically so we set it manually here
# We cant set vision attn same as general attn, because the general one can be sdpa if at
# least one sub-module (in this case LLM) supports sdpa
if hasattr(config, "vision_config"):
config.vision_config._attn_implementation = (
config._attn_implementation if config._attn_implementation != "sdpa" else "eager"
)
config.perceiver_config._attn_implementation = (
config._attn_implementation if config._attn_implementation != "sdpa" else "eager"
)
return config


Expand Down
8 changes: 5 additions & 3 deletions src/transformers/models/instructblip/modeling_instructblip.py
Original file line number Diff line number Diff line change
Expand Up @@ -1281,7 +1281,9 @@ class InstructBlipForConditionalGeneration(InstructBlipPreTrainedModel):
def __init__(self, config: InstructBlipConfig):
super().__init__(config)

self.vision_model = InstructBlipVisionModel(config.vision_config)
self.vision_model = InstructBlipVisionModel(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)

self.query_tokens = nn.Parameter(torch.zeros(1, config.num_query_tokens, config.qformer_config.hidden_size))
self.qformer = InstructBlipQFormerModel(config.qformer_config)
Expand All @@ -1290,11 +1292,11 @@ def __init__(self, config: InstructBlipConfig):

if config.use_decoder_only_language_model:
language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)
else:
language_model = AutoModelForSeq2SeqLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)

if language_model._no_split_modules is not None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1290,7 +1290,9 @@ class InstructBlipVideoForConditionalGeneration(InstructBlipVideoPreTrainedModel
def __init__(self, config: InstructBlipVideoConfig):
super().__init__(config)

self.vision_model = InstructBlipVideoVisionModel(config.vision_config)
self.vision_model = InstructBlipVideoVisionModel._from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)

self.query_tokens = nn.Parameter(torch.zeros(1, config.num_query_tokens, config.qformer_config.hidden_size))
self.qformer = InstructBlipVideoQFormerModel(config.qformer_config)
Expand All @@ -1299,11 +1301,11 @@ def __init__(self, config: InstructBlipVideoConfig):

if config.use_decoder_only_language_model:
language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)
else:
language_model = AutoModelForSeq2SeqLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)

if language_model._no_split_modules is not None:
Expand Down
14 changes: 4 additions & 10 deletions src/transformers/models/llava/modeling_llava.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,14 +149,6 @@ def _init_weights(self, module):
if module.padding_idx is not None:
module.weight.data[module.padding_idx].zero_()

@property
def _supports_sdpa(self):
"""
Retrieve language_model's attribute to check whether the model supports
SDPA or not.
"""
return self.language_model._supports_sdpa


LLAVA_INPUTS_DOCSTRING = r"""
Args:
Expand Down Expand Up @@ -236,12 +228,14 @@ def _supports_sdpa(self):
class LlavaForConditionalGeneration(LlavaPreTrainedModel):
def __init__(self, config: LlavaConfig):
super().__init__(config)
self.vision_tower = AutoModel.from_config(config.vision_config)
self.vision_tower = AutoModel.from_config(
config.vision_config, attn_implementation=config.text_config._attn_implementation
)

self.multi_modal_projector = LlavaMultiModalProjector(config)
self.vocab_size = config.text_config.vocab_size
self.language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.vision_config._attn_implementation
)
self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
self.post_init()
Expand Down
14 changes: 4 additions & 10 deletions src/transformers/models/llava_next/modeling_llava_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,14 +255,6 @@ def _init_weights(self, module):
if module.padding_idx is not None:
module.weight.data[module.padding_idx].zero_()

@property
def _supports_sdpa(self):
"""
Retrieve language_model's attribute to check whether the model supports
SDPA or not.
"""
return self.language_model._supports_sdpa


LLAVA_NEXT_INPUTS_DOCSTRING = r"""
Args:
Expand Down Expand Up @@ -345,15 +337,17 @@ def _supports_sdpa(self):
class LlavaNextForConditionalGeneration(LlavaNextPreTrainedModel):
def __init__(self, config: LlavaNextConfig):
super().__init__(config)
self.vision_tower = AutoModel.from_config(config.vision_config)
self.vision_tower = AutoModel.from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)

self.multi_modal_projector = LlavaNextMultiModalProjector(config)
embed_std = 1 / math.sqrt(config.text_config.hidden_size)
self.image_newline = nn.Parameter(torch.randn(config.text_config.hidden_size, dtype=self.dtype) * embed_std)

self.vocab_size = config.text_config.vocab_size
self.language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)
self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
self._padding_side = "left" # set it to left by default, user can use setter to change padding_sides
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -295,14 +295,6 @@ def _init_weights(self, module):
if module.padding_idx is not None:
module.weight.data[module.padding_idx].zero_()

@property
def _supports_sdpa(self):
"""
Retrieve language_model's attribute to check whether the model supports
SDPA or not.
"""
return self.language_model._supports_sdpa


LLAVA_NEXT_VIDEO_INPUTS_DOCSTRING = r"""
Args:
Expand Down Expand Up @@ -388,15 +380,17 @@ def __init__(
config: LlavaNextVideoConfig,
):
super().__init__(config)
self.vision_tower = AutoModel.from_config(config.vision_config)

self.vision_tower = AutoModel.from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)
self.multi_modal_projector = LlavaNextVideoMultiModalProjector(config)

embed_std = 1 / math.sqrt(config.text_config.hidden_size)
self.image_newline = nn.Parameter(torch.randn(config.text_config.hidden_size, dtype=self.dtype) * embed_std)

self.vocab_size = config.text_config.vocab_size
self.language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)
self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
self._padding_side = "left" # set it to left by default, user can use setter to change padding_sides
Expand Down
15 changes: 4 additions & 11 deletions src/transformers/models/paligemma/modeling_paligemma.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,6 @@ class PaliGemmaPreTrainedModel(PreTrainedModel):
_no_split_modules = ["PaliGemmaMultiModalProjector"]
_skip_keys_device_placement = "past_key_values"
_supports_flash_attn_2 = False
_supports_sdpa = True

def _init_weights(self, module):
# important: this ported version of PaliGemmaisn't meant for training from scratch - only
Expand All @@ -149,14 +148,6 @@ def _init_weights(self, module):
if module.padding_idx is not None:
module.weight.data[module.padding_idx].zero_()

@property
def _supports_sdpa(self):
"""
Retrieve language_model's attribute to check whether the model supports
SDPA or not.
"""
return self.language_model._supports_sdpa


PALIGEMMA_INPUTS_DOCSTRING = r"""
Args:
Expand Down Expand Up @@ -231,13 +222,15 @@ def _supports_sdpa(self):
class PaliGemmaForConditionalGeneration(PaliGemmaPreTrainedModel):
def __init__(self, config: PaliGemmaConfig):
super().__init__(config)
self.vision_tower = AutoModel.from_config(config=config.vision_config)
self.vision_tower = AutoModel.from_config(
config=config.vision_config, attn_implementation=config.vision_config._attn_implementation
)
self.multi_modal_projector = PaliGemmaMultiModalProjector(config)
self.vocab_size = config.text_config.vocab_size
self._attn_implementation = config._attn_implementation

language_model = AutoModelForCausalLM.from_config(
config=config.text_config, attn_implementation=self._attn_implementation
config=config.text_config, attn_implementation=config.text_config._attn_implementation
)

if language_model._tied_weights_keys is not None:
Expand Down
10 changes: 7 additions & 3 deletions src/transformers/models/video_llava/modeling_video_llava.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,13 +237,17 @@ def _supports_sdpa(self):
class VideoLlavaForConditionalGeneration(VideoLlavaPreTrainedModel):
def __init__(self, config: VideoLlavaConfig):
super().__init__(config)
self.video_tower = AutoModel.from_config(config.vision_config)
self.image_tower = AutoModel.from_config(config.vision_config)
self.video_tower = AutoModel.from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)
self.image_tower = AutoModel.from_config(
config.vision_config, attn_implementation=config.vision_config._attn_implementation
)

self.multi_modal_projector = VideoLlavaMultiModalProjector(config)
self.vocab_size = config.text_config.vocab_size
self.language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.text_config._attn_implementation
)
self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
self.post_init()
Expand Down
14 changes: 4 additions & 10 deletions src/transformers/models/vipllava/modeling_vipllava.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,14 +158,6 @@ def _init_weights(self, module):
if module.padding_idx is not None:
module.weight.data[module.padding_idx].zero_()

@property
def _supports_sdpa(self):
"""
Retrieve language_model's attribute to check whether the model supports
SDPA or not.
"""
return self.language_model._supports_sdpa


VIPLLAVA_INPUTS_DOCSTRING = r"""
Args:
Expand Down Expand Up @@ -241,12 +233,14 @@ def _supports_sdpa(self):
class VipLlavaForConditionalGeneration(VipLlavaPreTrainedModel):
def __init__(self, config: VipLlavaConfig):
super().__init__(config)
self.vision_tower = AutoModel.from_config(config.vision_config)
self.vision_tower = AutoModel.from_config(
config.vision_config, attn_implementation=config.text_config._attn_implementation
)

self.multi_modal_projector = VipLlavaMultiModalProjector(config)
self.vocab_size = config.text_config.vocab_size
self.language_model = AutoModelForCausalLM.from_config(
config.text_config, attn_implementation=config._attn_implementation
config.text_config, attn_implementation=config.vision_config._attn_implementation
)
self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
self.post_init()
Expand Down
Loading