Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
20c5c90043 | ||
|
|
c25bb86d1b | ||
|
|
89a0291cce | ||
|
|
f13f5d793e | ||
|
|
96fad37fc5 | ||
|
|
de5cc6cd40 |
@@ -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
|
||||
Reference in New Issue
Block a user