Compare commits

...
2 changed files with 124 additions and 3 deletions
+28 -3
View File
@@ -75,6 +75,25 @@ if not is_flash_attn_2_available():
logger = logging.get_logger(__name__)
def _resolve_torch_dtype(dtype, default: torch.dtype = torch.float32) -> torch.dtype:
if isinstance(dtype, torch.dtype):
return dtype
if isinstance(dtype, str):
normalized = dtype.strip().removeprefix("torch.").lower()
return {
"bfloat16": torch.bfloat16,
"bf16": torch.bfloat16,
"float16": torch.float16,
"fp16": torch.float16,
"half": torch.float16,
"float32": torch.float32,
"fp32": torch.float32,
"float": torch.float32,
}.get(normalized, default)
return default
class Qwen2_5_VLMLP(nn.Module):
def __init__(self, config, bias: bool = False):
super().__init__()
@@ -304,10 +323,13 @@ class Qwen2_5_VisionTransformerPretrainedModel(nn.Module):
config_class = Qwen2_5_VLVisionConfig
_no_split_modules = ["Qwen2_5_VLVisionBlock"]
def __init__(self, config) -> None:
def __init__(self, config, parent_torch_dtype=None) -> None:
super().__init__()
self.dtype = torch.bfloat16 if config.torch_dtype == "bfloat16" else torch.float32
config_torch_dtype = getattr(config, "torch_dtype", None)
self.dtype = _resolve_torch_dtype(
config_torch_dtype if config_torch_dtype is not None else parent_torch_dtype
)
self.spatial_merge_size = config.spatial_merge_size
self.patch_size = config.patch_size
@@ -1458,7 +1480,10 @@ class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
super().__init__()
config = _flatten_text_config(config)
self.config = config
self.visual = Qwen2_5_VisionTransformerPretrainedModel(config.vision_config)
self.visual = Qwen2_5_VisionTransformerPretrainedModel(
config.vision_config,
parent_torch_dtype=getattr(config, "torch_dtype", None),
)
self.model = Qwen2_5_VLModel(config)
self.vocab_size = config.vocab_size
@@ -0,0 +1,96 @@
# SPDX-License-Identifier: Apache-2.0
from types import SimpleNamespace
import torch
from fastvideo.models.encoders.qwen2_5_vl_custom import (
Qwen2_5_VisionTransformerPretrainedModel,
Qwen2_5_VLForConditionalGenerationSimple,
)
def _vision_config(torch_dtype=None):
return SimpleNamespace(
torch_dtype=torch_dtype,
spatial_merge_size=2,
patch_size=14,
fullatt_block_indexes=[],
window_size=112,
temporal_patch_size=2,
in_channels=3,
hidden_size=8,
num_heads=2,
depth=0,
_attn_implementation="sdpa",
out_hidden_size=8,
)
def _full_config(torch_dtype):
return SimpleNamespace(
vision_config=_vision_config(),
torch_dtype=torch_dtype,
vocab_size=16,
hidden_size=8,
num_hidden_layers=0,
num_attention_heads=2,
pad_token_id=0,
_attn_implementation="sdpa",
rms_norm_eps=1e-6,
max_position_embeddings=32,
rope_scaling=None,
)
def test_vision_dtype_uses_parent_bfloat16_string_when_vision_dtype_missing():
model = Qwen2_5_VisionTransformerPretrainedModel(
_vision_config(),
parent_torch_dtype="bfloat16",
)
assert model.dtype == torch.bfloat16
def test_vision_dtype_accepts_parent_torch_dtype_object():
model = Qwen2_5_VisionTransformerPretrainedModel(
_vision_config(),
parent_torch_dtype=torch.bfloat16,
)
assert model.dtype == torch.bfloat16
def test_vision_dtype_strips_parent_dtype_string_whitespace():
model = Qwen2_5_VisionTransformerPretrainedModel(
_vision_config(),
parent_torch_dtype=" torch.bfloat16 ",
)
assert model.dtype == torch.bfloat16
def test_vision_dtype_prefers_explicit_vision_dtype_over_parent_dtype():
model = Qwen2_5_VisionTransformerPretrainedModel(
_vision_config(torch_dtype="float16"),
parent_torch_dtype="bfloat16",
)
assert model.dtype == torch.float16
def test_vision_dtype_falls_back_to_float32_for_missing_or_unknown_dtype():
missing = Qwen2_5_VisionTransformerPretrainedModel(_vision_config())
unknown = Qwen2_5_VisionTransformerPretrainedModel(
_vision_config(),
parent_torch_dtype="not-a-real-dtype",
)
assert missing.dtype == torch.float32
assert unknown.dtype == torch.float32
def test_conditional_generation_passes_parent_dtype_to_visual_tower():
model = Qwen2_5_VLForConditionalGenerationSimple(_full_config("torch.bfloat16"))
assert model.visual.dtype == torch.bfloat16