2828
2929# Global cache of HPSv3 inferencers keyed by device.
3030_HPSV3_INFERENCERS : dict [str , Any ] = {}
31+ _HPSV3_LOAD_PATCHED = False
32+
33+
34+ def _patch_transformers_video_input_alias () -> None :
35+ """Keep HPSv3 compatible with newer transformers releases.
36+
37+ HPSv3 imports ``VideoInput`` from ``transformers.image_utils`` for type
38+ annotations. Some transformers versions used by FastVideo no longer
39+ export that alias, even though the runtime image utilities HPSv3 needs are
40+ still present.
41+ """
42+ from transformers import image_utils
43+
44+ if not hasattr (image_utils , "VideoInput" ):
45+ image_utils .VideoInput = image_utils .ImageInput
46+
47+
48+ def _remap_hpsv3_state_dict (state_dict : dict [str , Any ]) -> dict [str , Any ]:
49+ """Adapt HPSv3 checkpoints saved with older Qwen2-VL key names."""
50+ remapped = {}
51+ for key , value in state_dict .items ():
52+ if key .startswith ("visual." ):
53+ key = f"model.{ key } "
54+ elif key .startswith ("model.layers." ):
55+ key = f"model.language_model.{ key [len ('model.' ):]} "
56+ elif key .startswith ("model.embed_tokens." ):
57+ key = f"model.language_model.{ key [len ('model.' ):]} "
58+ elif key .startswith ("model.norm." ):
59+ key = f"model.language_model.{ key [len ('model.' ):]} "
60+
61+ key = key .replace (
62+ "base_model.model.visual." ,
63+ "base_model.model.model.visual." ,
64+ 1 ,
65+ )
66+ key = key .replace (
67+ "base_model.model.model.layers." ,
68+ "base_model.model.model.language_model.layers." ,
69+ 1 ,
70+ )
71+ key = key .replace (
72+ "base_model.model.model.embed_tokens." ,
73+ "base_model.model.model.language_model.embed_tokens." ,
74+ 1 ,
75+ )
76+ key = key .replace (
77+ "base_model.model.model.norm." ,
78+ "base_model.model.model.language_model.norm." ,
79+ 1 ,
80+ )
81+ remapped [key ] = value
82+ return remapped
83+
84+
85+ def _walk_model_graph (model : Any ):
86+ """Yield common wrapper/base model objects without importing PEFT."""
87+ stack = [model ]
88+ seen = set ()
89+ while stack :
90+ current = stack .pop ()
91+ if current is None or id (current ) in seen :
92+ continue
93+ seen .add (id (current ))
94+ yield current
95+ for attr in ("base_model" , "model" ):
96+ child = getattr (current , attr , None )
97+ if child is not None :
98+ stack .append (child )
99+
100+
101+ def _patch_load_state_dict (cls : Any ) -> None :
102+ """Patch a model class to accept old Qwen2-VL checkpoint keys."""
103+ if getattr (cls , "_fastvideo_qwen2vl_key_remap" , False ):
104+ return
105+
106+ original_load_state_dict = cls .load_state_dict
107+
108+ def load_state_dict_with_key_remap (
109+ self ,
110+ state_dict ,
111+ strict = True ,
112+ assign = False ,
113+ ):
114+ state_dict = _remap_hpsv3_state_dict (state_dict )
115+ return original_load_state_dict (
116+ self ,
117+ state_dict ,
118+ strict = strict ,
119+ assign = assign ,
120+ )
121+
122+ cls .load_state_dict = load_state_dict_with_key_remap
123+ cls ._fastvideo_qwen2vl_key_remap = True
124+
125+
126+ def _patch_hpsv3_state_dict_loader () -> None :
127+ """Patch HPSv3 reward model loading for transformers key drift."""
128+ global _HPSV3_LOAD_PATCHED
129+ if _HPSV3_LOAD_PATCHED :
130+ return
131+
132+ from hpsv3 .model .qwen2vl_trainer import Qwen2VLRewardModelBT
133+
134+ _patch_load_state_dict (Qwen2VLRewardModelBT )
135+ try :
136+ from peft import PeftModel
137+ except ImportError :
138+ PeftModel = None
139+ if PeftModel is not None :
140+ _patch_load_state_dict (PeftModel )
141+ _HPSV3_LOAD_PATCHED = True
142+
143+
144+ def _patch_hpsv3_runtime_model (model : Any ) -> None :
145+ """Add aliases expected by HPSv3's older Qwen2-VL forward."""
146+ for candidate in _walk_model_graph (model ):
147+ language_model = getattr (candidate , "language_model" , None )
148+ if (
149+ language_model is not None
150+ and not hasattr (candidate , "embed_tokens" )
151+ and hasattr (language_model , "embed_tokens" )
152+ ):
153+ candidate .__dict__ ["embed_tokens" ] = language_model .embed_tokens
31154
32155
33156def _normalize_device (device ) -> str :
@@ -36,6 +159,19 @@ def _normalize_device(device) -> str:
36159 return str (torch .device (device ))
37160
38161
162+ def _move_hpsv3_inferencer (inferencer : Any , device ) -> None :
163+ """Move an HPSv3 inferencer across devices.
164+
165+ HPSv3RewardInferencer does not expose ``.to()``, but it stores its torch
166+ module on ``.model`` and reads ``.device`` when preparing batches.
167+ """
168+ device_str = _normalize_device (device )
169+ model = getattr (inferencer , "model" , None )
170+ if model is not None and hasattr (model , "to" ):
171+ model .to (device )
172+ inferencer .device = device_str
173+
174+
39175def set_hpsv3_device (device ) -> None :
40176 """Move cached HPSv3 inferencer to given device."""
41177 key = _normalize_device (device )
@@ -44,7 +180,7 @@ def set_hpsv3_device(device) -> None:
44180 # Move from any existing device.
45181 for old_key , inf in list (_HPSV3_INFERENCERS .items ()):
46182 if old_key != key :
47- inf . to ( device )
183+ _move_hpsv3_inferencer ( inf , device )
48184 _HPSV3_INFERENCERS [key ] = inf
49185 del _HPSV3_INFERENCERS [old_key ]
50186 return
@@ -55,15 +191,18 @@ def _get_hpsv3_inferencer(device):
55191 key = _normalize_device (device )
56192 if key not in _HPSV3_INFERENCERS :
57193 try :
194+ _patch_transformers_video_input_alias ()
58195 from hpsv3 import HPSv3RewardInferencer
196+ _patch_hpsv3_state_dict_loader ()
59197 except ImportError as exc :
60198 msg = (
61- "hpsv3 package not found . Ensure the "
62- "HPSv3 submodule is checked out under "
63- "fastvideo/train/methods/rl/reward/HPSv3 "
199+ "Failed to import HPSv3 . Ensure the HPSv3 submodule is "
200+ "checked out under fastvideo/train/methods/rl/reward/HPSv3 "
201+ "and that its transformers dependencies are compatible. "
64202 )
65203 raise ImportError (msg ) from exc
66204 inf = HPSv3RewardInferencer (device = device )
205+ _patch_hpsv3_runtime_model (inf .model )
67206 _HPSV3_INFERENCERS [key ] = inf
68207 return _HPSV3_INFERENCERS [key ]
69208
0 commit comments