Skip to content

Commit 28268ce

Browse files
AbecidDavids048
andauthored
[feat] GenRL: stabilize reward model compatibility (#1400)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
1 parent 78cbbeb commit 28268ce

5 files changed

Lines changed: 471 additions & 16 deletions

File tree

GenRL

Lines changed: 0 additions & 1 deletion
This file was deleted.

fastvideo/train/methods/rl/reward/hpsv3.py

Lines changed: 143 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,129 @@
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

33156
def _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+
39175
def 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

fastvideo/train/methods/rl/reward/utils.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,11 +24,20 @@ def prepare_images(
2424

2525
if images.ndim == 4:
2626
# Image batch: (N, C, H, W) or (N, H, W, C)
27-
if images.shape[1] in (1, 3):
27+
if images.shape[-1] in (1, 3):
28+
pass
29+
elif images.shape[1] in (1, 3):
2830
images = images.transpose(0, 2, 3, 1)
2931
elif images.ndim == 5:
30-
# Video batch: (N, F, C, H, W) or (N, C, F, H, W)
31-
if images.shape[2] in (1, 3):
32+
# Video batch: (N, F, H, W, C), (N, F, C, H, W),
33+
# or (N, C, F, H, W). Check channel-last first because
34+
# one-frame videos have shape[1] == 1.
35+
if images.shape[-1] in (1, 3):
36+
pass
37+
elif images.shape[1] == 3 and images.shape[2] == 1:
38+
# (N, C=3, F=1, H, W) -> (N, F, H, W, C)
39+
images = images.transpose(0, 2, 3, 4, 1)
40+
elif images.shape[2] in (1, 3):
3241
# (N, F, C, H, W) -> (N, F, H, W, C)
3342
images = images.transpose(0, 1, 3, 4, 2)
3443
elif images.shape[1] in (1, 3):

0 commit comments

Comments
 (0)