mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
fix(annotation): skip vlm initialization when using advantage module
This commit is contained in:
@@ -186,7 +186,7 @@ class Executor:
|
|||||||
staging_dir: Path,
|
staging_dir: Path,
|
||||||
module: Any,
|
module: Any,
|
||||||
) -> PhaseResult:
|
) -> PhaseResult:
|
||||||
if not module.enabled:
|
if module is None or not module.enabled:
|
||||||
print(f"[annotate] phase={name} skipped (module disabled)", flush=True)
|
print(f"[annotate] phase={name} skipped (module disabled)", flush=True)
|
||||||
return PhaseResult(name=name, episodes_processed=0, episodes_skipped=len(records))
|
return PhaseResult(name=name, episodes_processed=0, episodes_skipped=len(records))
|
||||||
n = len(records)
|
n = len(records)
|
||||||
@@ -238,7 +238,7 @@ class Executor:
|
|||||||
``plan`` module with the interjection timestamps so its existing
|
``plan`` module with the interjection timestamps so its existing
|
||||||
prompt path is reused.
|
prompt path is reused.
|
||||||
"""
|
"""
|
||||||
if not self.plan.enabled or not self.interjections.enabled:
|
if not self.plan or not self.plan.enabled or not self.interjections or not self.interjections.enabled:
|
||||||
return PhaseResult(name="plan_update", episodes_processed=0, episodes_skipped=len(records))
|
return PhaseResult(name="plan_update", episodes_processed=0, episodes_skipped=len(records))
|
||||||
processed = 0
|
processed = 0
|
||||||
for record in records:
|
for record in records:
|
||||||
|
|||||||
@@ -64,12 +64,18 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
|
|||||||
root = _resolve_root(cfg)
|
root = _resolve_root(cfg)
|
||||||
logger.info("annotate: root=%s", root)
|
logger.info("annotate: root=%s", root)
|
||||||
|
|
||||||
vlm = make_vlm_client(cfg.vlm)
|
needs_vlm = cfg.plan.enabled or cfg.interjections.enabled or cfg.vqa.enabled
|
||||||
frame_provider = make_frame_provider(root, camera_key=cfg.vlm.camera_key, video_backend=cfg.video_backend)
|
vlm = make_vlm_client(cfg.vlm) if needs_vlm else None
|
||||||
|
frame_provider = (
|
||||||
|
make_frame_provider(root, camera_key=cfg.vlm.camera_key, video_backend=cfg.video_backend)
|
||||||
|
if needs_vlm
|
||||||
|
else None
|
||||||
|
)
|
||||||
# Surface the resolved cameras up front so a silent vqa-module no-op
|
# Surface the resolved cameras up front so a silent vqa-module no-op
|
||||||
# is obvious in job output rather than discovered post-hoc by counting
|
# is obvious in job output rather than discovered post-hoc by counting
|
||||||
# parquet rows.
|
# parquet rows.
|
||||||
cam_keys = list(getattr(frame_provider, "camera_keys", []) or [])
|
cam_keys = list(getattr(frame_provider, "camera_keys", []) or []) if frame_provider else []
|
||||||
|
if frame_provider:
|
||||||
logger.info(
|
logger.info(
|
||||||
"annotate: frame_provider default camera=%r, all cameras=%s",
|
"annotate: frame_provider default camera=%r, all cameras=%s",
|
||||||
getattr(frame_provider, "camera_key", None),
|
getattr(frame_provider, "camera_key", None),
|
||||||
@@ -82,15 +88,27 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
|
|||||||
"meta/info.json for observation.images.* features, or pass "
|
"meta/info.json for observation.images.* features, or pass "
|
||||||
"--vlm.camera_key=<key> to seed the cameras list."
|
"--vlm.camera_key=<key> to seed the cameras list."
|
||||||
)
|
)
|
||||||
plan = PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan, frame_provider=frame_provider)
|
plan = (
|
||||||
interjections = InterjectionsAndSpeechModule(
|
PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan, frame_provider=frame_provider)
|
||||||
|
if needs_vlm
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
interjections = (
|
||||||
|
InterjectionsAndSpeechModule(
|
||||||
vlm=vlm, config=cfg.interjections, seed=cfg.seed, frame_provider=frame_provider
|
vlm=vlm, config=cfg.interjections, seed=cfg.seed, frame_provider=frame_provider
|
||||||
)
|
)
|
||||||
vqa = GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed, frame_provider=frame_provider)
|
if needs_vlm
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
vqa = (
|
||||||
|
GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed, frame_provider=frame_provider)
|
||||||
|
if needs_vlm
|
||||||
|
else None
|
||||||
|
)
|
||||||
advantage = AdvantageModule(config=cfg.advantage)
|
advantage = AdvantageModule(config=cfg.advantage)
|
||||||
writer = LanguageColumnsWriter()
|
writer = LanguageColumnsWriter()
|
||||||
validator = StagingValidator(
|
validator = StagingValidator(
|
||||||
dataset_camera_keys=tuple(getattr(frame_provider, "camera_keys", []) or []) or None,
|
dataset_camera_keys=tuple(cam_keys) or None,
|
||||||
)
|
)
|
||||||
|
|
||||||
executor = Executor(
|
executor = Executor(
|
||||||
|
|||||||
Reference in New Issue
Block a user