unsloth/tests/test_video_path_validation.py
Leo Borcherding 756e129388
fix/validate dataset video paths before training (#5136)
* fix(vision): validate dataset video paths before training

* fix(vision): remove redundant warnings import, add pytest tests for #5085

* fix(trainer): auto-validate video paths in UnslothVisionDataCollator on first batch (#5085)

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* refactor(vision): use str.removeprefix instead of slicing (Datta0 review)

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* style: replace em dashes with hyphens in error message and docstring

* Update unsloth/models/vision.py

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

* Split: keep only 1 file(s)

* fix(vision): broaden video-path validation to all collator inputs

- check_dataset_for_missing_videos now accepts every example shape that
  UnslothVisionDataCollator forwards to process_vision_info:
  dict rows with "messages"/"conversations"/"prompt"/"completion", and raw
  message-list rows.  Earlier logic only handled {"messages": [...]}, so
  conversations/prompt/completion datasets silently skipped validation and
  raw message-list rows crashed on list.get.
- Guard against non-dict message entries; a bare string inside a message
  list no longer raises AttributeError.
- Decode file:// URIs via urllib so percent-encoded paths, absolute
  Windows URIs (file:///C:/...) and host-qualified URIs
  (file://localhost/abs/path) map back to their real filesystem path.
- Expose an optional "checked" set so callers can reuse dedup state across
  invocations.
- Docstring warns that passing a streaming IterableDataset consumes the
  iterator.

UnslothVisionDataCollator now validates every batch (not just batch 0)
and applies formatting_func before validation, matching the base
collator's own ordering so formatter-generated video paths are also
checked.  The already-checked set is shared across batches, so per-batch
cost stays proportional to newly seen paths.

* fix(vision): robust URI + scheme handling in video-path validator

- _local_path_from_video_value now treats anything with a "://" prefix
  as a URI and validates only file:// URIs.  This prevents false
  FileNotFoundError on remote schemes that were silently passed through
  before (s3://, gs://, hf://, ftp://, az://, ...).
- Non-localhost file authorities (e.g. file://nas-server/share/clip.mp4)
  are now skipped instead of being stripped and validated against the
  local filesystem; RFC 8089 only permits empty host or "localhost" for
  local files.
- Drop the explicit unquote call: urllib.request.url2pathname already
  unquotes, so the previous url2pathname(unquote(path)) double-decoded
  any filename with a literal percent (e.g. a file named "clip%20.mp4").
- Remove the Windows drive-letter strip block; nturl2path.url2pathname
  handles "/C:/foo" -> "C:\\foo" itself, leaving nothing for the guard to
  match on either OS.
- Return None when the resolved path is empty (bare "file://" or
  "file://hostname") so the caller skips it instead of reporting a
  blank "  - " entry in the error message.
- Add a runtime guard in check_dataset_for_missing_videos that warns and
  returns early when handed a datasets.IterableDataset, matching the
  docstring contract and preventing silent iterator exhaustion.

Windows native paths like "C:/path/x.mp4" stay valid because the
scheme check uses the "://" substring (not urlparse's single-letter
scheme surface).

* tests: consolidate video-path validation coverage into one file

New coverage for tests/test_video_path_validation.py:

- every-batch validation with cross-batch dedup (replaces the old
  first-batch-only assertion which no longer matches the implementation).
- all collator-supported input shapes: messages, conversations,
  prompt/completion, raw-message-list rows; non-dict message entries.
- file:// URI robustness: percent-encoded paths, localhost netloc,
  non-localhost netloc skipped, bare / hostname-only URIs skipped,
  double-encoded filenames single-unquote correctly.
- non-file remote schemes (s3, gs, hf, ftp, az) skipped without raising.
- Windows-style absolute path not mistaken for a URI scheme.
- formatting_func applied before validation inside the collator wrapper.
- IterableDataset runtime guard warns and returns without consuming.

The pre-existing test_collator_validates_only_once assertion has been
replaced by test_collator_validates_every_batch / dedupes_across_batches
because the wrapper now validates every batch.

The session-scoped AST fallback fixture was extended to extract the
helper functions that the rewritten check_dataset_for_missing_videos
depends on, so the Windows/no-triton code path still loads the module
surface.

* Fix CI iter 1

* Omit release-desktop.yml from the PR diff

The workflow file landed on origin/main after the PR branched off; our
fork-scoped push token cannot touch .github/workflows/**. Drop it from
this branch so the PR diff stays within the author's authorisable
surface. The file remains on origin/main and will return after this PR
merges upstream.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* fix(vision): cache only validated paths in checked set

Missing paths were added to the dedup cache before the existence check,
so a caller that caught FileNotFoundError and retried with the same
collator/checked set would silently skip the bad path on the second call.

Only add a path to checked after os.path.isfile confirms it exists, so
missing paths are re-validated on every call until they are fixed.

* fix(ci): resolve two CI failures introduced by this PR

- Add __all__ to models/__init__.py so the HOISTED-IMPORT-UNUSED linter
  check passes for check_dataset_for_missing_videos
- Replace Dataset.from_list() in test helpers with plain list literals;
  the CI environment mocks datasets with a MockDataset that only has
  from_dict, but check_dataset_for_missing_videos accepts any iterable
  so no Dataset wrapper is needed
- Guard test_iterable_dataset_warns_and_skips with pytest.importorskip
  so it skips cleanly when the real datasets package is unavailable

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* fix(ci): correct HOISTED-IMPORT-UNUSED without breaking wildcard exports

Adding __all__ to models/__init__.py was too aggressive - it restricted
from .models import * in _gpu_init.py to only check_dataset_for_missing_videos,
hiding FastLanguageModel, FastVisionModel etc and breaking
test_fast_model_class_surface_under_spoof.

Instead: remove __all__, and add an explicit named import in _gpu_init.py
so the linter sees the symbol consumed in the re-export chain.

* Deduplicate missing video paths and skip data URIs in validator

check_dataset_for_missing_videos appended a path to the missing list on every
occurrence, so a path referenced by multiple rows was reported N times and the
error header read the wrong count. Track missing paths in a per-call set so each
is listed once, kept separate from the checked cache so retries still re-check
missing files. This restores the dedup behaviour the docstring promises and the
existing test_duplicate_paths_deduplicated test asserts.

Also skip data: URIs in _local_path_from_video_value so inline base64 payloads
are not flagged as missing files.

Add tests for the data URI case, warn-only dedup, and real integration against
the unsloth_zoo UnslothVisionDataCollator base (verifying validation gates the
base call, formatting_func is applied once and restored even when the base
raises).

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* fix(trainer): declare __slots__ on UnslothVisionDataCollator subclass

* fix(ci): hoist check_dataset_for_missing_videos to trainer module level; guard IterableDataset skip

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* fix(vision): accept tuple message content in video path validator

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Tighten comments and docstrings for PR #5136

Comment-only pass over the new video path validation code: shorten the
collator and validator docstrings, collapse multi-line inline comments,
and reduce test docstrings to one-liners. No code changes; verified with
comment_tools.py check --strip-docstrings (3/3 code unchanged) and the
full test suite (35 passed).

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
2026-06-10 05:32:32 -07:00

570 lines
20 KiB
Python

"""
Tests for check_dataset_for_missing_videos (issue #5085).
Fixtures extract the function from vision.py via AST so the pure-Python logic
tests run without the full unsloth import chain (triton/CUDA kernels).
"""
import ast
import os
import tempfile
import warnings
from pathlib import Path
import pytest
# ── Fixtures ──────────────────────────────────────────────────────────────────
def _extract_fns_via_ast(
source_path,
fn_names,
extra_ns = None,
):
"""Exec a set of top-level functions out of a .py file so intra-module
references between them resolve."""
source = source_path.read_text(encoding = "utf-8")
tree = ast.parse(source, filename = str(source_path))
wanted = set(fn_names)
nodes = [n for n in tree.body if isinstance(n, ast.FunctionDef) and n.name in wanted]
missing = wanted - {n.name for n in nodes}
if missing:
pytest.fail(f"{sorted(missing)} not found in {source_path}")
mini = ast.Module(body = nodes, type_ignores = [])
ast.fix_missing_locations(mini)
ns = {"os": os, "warnings": warnings, "__name__": "_extracted"}
if extra_ns:
ns.update(extra_ns)
exec(compile(mini, str(source_path), "exec"), ns)
return {name: ns[name] for name in fn_names}
def _extract_fn_via_ast(
source_path,
fn_name,
extra_ns = None,
):
return _extract_fns_via_ast(source_path, [fn_name], extra_ns)[fn_name]
@pytest.fixture(scope = "session")
def check_dataset_for_missing_videos():
"""Direct import when possible, else AST extraction from vision.py."""
try:
from unsloth.models.vision import check_dataset_for_missing_videos as fn
return fn
except Exception:
pass
vision_path = Path(__file__).parent.parent / "unsloth" / "models" / "vision.py"
fns = _extract_fns_via_ast(
vision_path,
[
"_looks_like_message_list",
"_iter_message_lists",
"_local_path_from_video_value",
"check_dataset_for_missing_videos",
],
)
return fns["check_dataset_for_missing_videos"]
@pytest.fixture(scope = "session")
def make_auto_validating_collator(check_dataset_for_missing_videos):
"""Factory for a minimal collator mirroring the trainer.py wrapper."""
class _FakeBase:
def __init__(self, formatting_func = None):
self.formatting_func = formatting_func
def __call__(self, examples):
if self.formatting_func is not None:
examples = [self.formatting_func(e) for e in examples]
return {"ok": True, "examples": examples}
class _AutoValidatingCollator(_FakeBase):
def __init__(self, formatting_func = None):
super().__init__(formatting_func = formatting_func)
self._checked_video_paths = set()
def __call__(self, examples):
formatting_func = self.formatting_func
if formatting_func is not None:
examples = [formatting_func(e) for e in examples]
check_dataset_for_missing_videos(
examples,
raise_error = True,
checked = self._checked_video_paths,
)
if formatting_func is None:
return super().__call__(examples)
self.formatting_func = None
try:
return super().__call__(examples)
finally:
self.formatting_func = formatting_func
return _AutoValidatingCollator
# ── Helpers ───────────────────────────────────────────────────────────────────
def _make_video_dataset(*video_paths):
return [
{"messages": [{"role": "user", "content": [{"type": "video", "video": p}]}]}
for p in video_paths
]
def _batch(*video_paths):
return _make_video_dataset(*video_paths)
# ── Tests: check_dataset_for_missing_videos ───────────────────────────────────
def test_missing_local_file_raises(check_dataset_for_missing_videos):
"""Missing local path raises FileNotFoundError."""
ds = _make_video_dataset("/nonexistent/videos/clip.mp4")
with pytest.raises(FileNotFoundError):
check_dataset_for_missing_videos(ds)
def test_remote_url_skipped(check_dataset_for_missing_videos):
"""http/https URLs are not checked locally."""
ds = _make_video_dataset("https://example.com/video.mp4")
assert check_dataset_for_missing_videos(ds) == []
def test_existing_file_accepted(check_dataset_for_missing_videos):
"""Existing local file passes without error."""
with tempfile.NamedTemporaryFile(suffix = ".mp4", delete = False) as f:
f.write(b"fake video bytes")
tmp = f.name
try:
ds = _make_video_dataset(tmp)
assert check_dataset_for_missing_videos(ds) == []
finally:
os.unlink(tmp)
def test_file_uri_scheme_stripped(check_dataset_for_missing_videos):
"""file:// scheme is stripped before the path check."""
ds = _make_video_dataset("file:///nonexistent/clip.mp4")
with pytest.raises(FileNotFoundError):
check_dataset_for_missing_videos(ds)
def test_warn_only_mode(check_dataset_for_missing_videos):
"""raise_error=False warns and returns the missing paths."""
ds = _make_video_dataset("/nonexistent/videos/clip.mp4")
with warnings.catch_warnings(record = True) as caught:
warnings.simplefilter("always")
missing = check_dataset_for_missing_videos(ds, raise_error = False)
assert len(caught) == 1
assert "could not be found" in str(caught[0].message)
assert missing == ["/nonexistent/videos/clip.mp4"]
def test_duplicate_paths_deduplicated(check_dataset_for_missing_videos):
"""Repeated missing path is listed once."""
ds = _make_video_dataset("/nonexistent/clip.mp4", "/nonexistent/clip.mp4")
with pytest.raises(FileNotFoundError) as exc_info:
check_dataset_for_missing_videos(ds)
assert str(exc_info.value).count("/nonexistent/clip.mp4") == 1
# ── Tests: UnslothVisionDataCollator auto-validation ─────────────────────────
def test_collator_raises_on_first_batch_with_missing_video(make_auto_validating_collator):
"""Collator raises on a missing path with no user action needed."""
collator = make_auto_validating_collator()
batch = _batch("/nonexistent/auto/clip.mp4")
with pytest.raises(FileNotFoundError):
collator(batch)
def test_collator_passes_on_first_batch_with_valid_video(make_auto_validating_collator):
"""Collator passes a valid batch through."""
with tempfile.NamedTemporaryFile(suffix = ".mp4", delete = False) as f:
f.write(b"fake video bytes")
tmp = f.name
try:
collator = make_auto_validating_collator()
batch = _batch(tmp)
result = collator(batch)
assert result["ok"] is True
finally:
os.unlink(tmp)
def test_collator_validates_every_batch(make_auto_validating_collator):
"""A missing video first appearing after batch 0 must still raise."""
with tempfile.NamedTemporaryFile(suffix = ".mp4", delete = False) as f:
f.write(b"fake video bytes")
tmp = f.name
try:
collator = make_auto_validating_collator()
collator(_batch(tmp)) # batch 0: valid
with pytest.raises(FileNotFoundError):
collator(_batch("/nonexistent/late.mp4"))
finally:
os.unlink(tmp)
def test_collator_dedupes_across_batches(make_auto_validating_collator):
"""The checked-path set is shared across batches."""
with tempfile.NamedTemporaryFile(suffix = ".mp4", delete = False) as f:
f.write(b"fake video bytes")
tmp = f.name
try:
collator = make_auto_validating_collator()
collator(_batch(tmp))
collator(_batch(tmp, tmp))
assert tmp in collator._checked_video_paths
finally:
os.unlink(tmp)
def test_conversations_column_missing_detected(check_dataset_for_missing_videos):
"""'conversations' column is scanned."""
ds = [
{
"conversations": [
{
"role": "user",
"content": [{"type": "video", "video": "/nonexistent/conv.mp4"}],
}
]
},
]
with pytest.raises(FileNotFoundError):
check_dataset_for_missing_videos(ds)
def test_prompt_completion_column_missing_detected(check_dataset_for_missing_videos):
"""'prompt'/'completion' columns are scanned."""
ds = [
{
"prompt": [
{
"role": "user",
"content": [{"type": "video", "video": "/nonexistent/p.mp4"}],
}
],
"completion": [{"role": "assistant", "content": [{"type": "text", "text": "hi"}]}],
},
]
with pytest.raises(FileNotFoundError) as exc_info:
check_dataset_for_missing_videos(ds)
assert "/nonexistent/p.mp4" in str(exc_info.value)
def test_raw_message_list_example_missing_detected(check_dataset_for_missing_videos):
"""Rows that are themselves message lists (no outer dict) are scanned."""
ds = [
[
{
"role": "user",
"content": [{"type": "video", "video": "/nonexistent/raw.mp4"}],
}
],
]
with pytest.raises(FileNotFoundError):
check_dataset_for_missing_videos(ds)
def test_non_dict_message_entry_does_not_crash(check_dataset_for_missing_videos):
"""Non-dict message entries are skipped."""
ds = [{"messages": ["not a dict", {"role": "user", "content": []}]}]
assert check_dataset_for_missing_videos(ds) == []
def test_file_uri_percent_encoded(check_dataset_for_missing_videos, tmp_path):
"""Percent-encoded file:// URIs decode to the real path."""
target = tmp_path / "my video.mp4"
target.write_bytes(b"x")
uri = "file://" + str(target).replace(" ", "%20")
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": uri}]}]}]
assert check_dataset_for_missing_videos(ds) == []
def test_file_uri_localhost_host(check_dataset_for_missing_videos, tmp_path):
"""file://localhost/<abs path> is the local machine (RFC 8089)."""
target = tmp_path / "clip.mp4"
target.write_bytes(b"x")
uri = f"file://localhost{target}"
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": uri}]}]}]
assert check_dataset_for_missing_videos(ds) == []
def test_checked_set_reused_across_calls(check_dataset_for_missing_videos, tmp_path):
"""A supplied checked set is populated and deduped across calls."""
target = tmp_path / "clip.mp4"
target.write_bytes(b"x")
shared = set()
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": str(target)}]}]}]
check_dataset_for_missing_videos(ds, checked = shared)
assert str(target) in shared
check_dataset_for_missing_videos(ds, checked = shared)
assert len(shared) == 1
@pytest.mark.parametrize(
"uri",
[
"s3://bucket/clip.mp4",
"gs://bucket/clip.mp4",
"hf://datasets/u/r/clip.mp4",
"ftp://host/clip.mp4",
"az://container/clip.mp4",
],
)
def test_non_file_remote_scheme_skipped(check_dataset_for_missing_videos, uri):
"""Non-file URI schemes are treated as remote and skipped."""
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": uri}]}]}]
assert check_dataset_for_missing_videos(ds) == []
def test_file_uri_non_localhost_host_skipped(check_dataset_for_missing_videos):
"""file://<non-localhost>/path is remote (RFC 8089): skip local checks."""
ds = [
{
"messages": [
{
"role": "user",
"content": [{"type": "video", "video": "file://nas-server/share/clip.mp4"}],
}
]
}
]
assert check_dataset_for_missing_videos(ds) == []
@pytest.mark.parametrize("uri", ["file://", "file://hostname"])
def test_degenerate_file_uri_skipped(check_dataset_for_missing_videos, uri):
"""No path component must not produce a blank missing entry."""
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": uri}]}]}]
assert check_dataset_for_missing_videos(ds) == []
def test_file_uri_double_encoded_percent(check_dataset_for_missing_videos, tmp_path):
"""%2520 must single-unquote to 'clip%20.mp4', not 'clip .mp4'."""
target = tmp_path / "clip%20.mp4"
target.write_bytes(b"x")
uri = "file://" + str(target).replace("%", "%25")
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": uri}]}]}]
assert check_dataset_for_missing_videos(ds) == []
def test_windows_style_absolute_path_not_mistaken_for_scheme(
check_dataset_for_missing_videos, tmp_path
):
"""'C:/...' has no '://' so it is a plain path, even where urlparse
would yield scheme='c'."""
target = tmp_path / "clip.mp4"
target.write_bytes(b"x")
path = str(target)
if os.name != "nt":
# '://'-free values must round-trip unchanged; keep the real path
path = str(target)
ds = [{"messages": [{"role": "user", "content": [{"type": "video", "video": path}]}]}]
assert check_dataset_for_missing_videos(ds) == []
ds_missing = [
{
"messages": [
{
"role": "user",
"content": [{"type": "video", "video": "C:/definitely/missing.mp4"}],
}
]
}
]
with pytest.raises(FileNotFoundError) as exc:
check_dataset_for_missing_videos(ds_missing)
assert "C:/definitely/missing.mp4" in str(exc.value)
def test_iterable_dataset_warns_and_skips(check_dataset_for_missing_videos):
"""Streaming IterableDataset: warn, return [], do not exhaust it."""
datasets_mod = pytest.importorskip("datasets", reason = "real datasets package required")
if not hasattr(datasets_mod, "IterableDataset"):
pytest.skip("datasets.IterableDataset not available in this environment")
IterableDataset = datasets_mod.IterableDataset
def gen():
for p in ("/nonexistent/a.mp4", "/nonexistent/b.mp4"):
yield {"messages": [{"role": "user", "content": [{"type": "video", "video": p}]}]}
ds = IterableDataset.from_generator(gen)
with warnings.catch_warnings(record = True) as caught:
warnings.simplefilter("always")
result = check_dataset_for_missing_videos(ds)
assert result == []
assert any("IterableDataset" in str(w.message) for w in caught)
# generator must not have been exhausted
consumed = list(ds)
assert len(consumed) == 2
def test_collator_applies_formatting_func_before_validation(make_auto_validating_collator):
"""formatting_func runs before validation; super gets formatted examples
and must not re-apply it."""
def fmt(example):
return {
"messages": [
{
"role": "user",
"content": [{"type": "video", "video": example["video_id"]}],
}
]
}
raise_collator = make_auto_validating_collator(formatting_func = fmt)
with pytest.raises(FileNotFoundError):
raise_collator([{"video_id": "/nonexistent/formatted.mp4"}])
with tempfile.NamedTemporaryFile(suffix = ".mp4", delete = False) as f:
f.write(b"x")
tmp = f.name
try:
ok_collator = make_auto_validating_collator(formatting_func = fmt)
before = ok_collator.formatting_func
result = ok_collator([{"video_id": tmp}])
assert result["ok"] is True
assert ok_collator.formatting_func is before
passed = result["examples"]
assert passed[0]["messages"][0]["content"][0]["video"] == tmp
finally:
os.unlink(tmp)
def test_data_uri_skipped(check_dataset_for_missing_videos):
"""Inline data: URIs are not flagged missing."""
ds = _make_video_dataset("data:video/mp4;base64,AAAABBBBCCCC")
assert check_dataset_for_missing_videos(ds) == []
def test_tuple_content_entries_checked(check_dataset_for_missing_videos):
"""Tuple message content is validated like a list."""
ds = [
{
"messages": [
{
"role": "user",
"content": ({"type": "video", "video": "/nonexistent/tuple.mp4"},),
}
]
}
]
with pytest.raises(FileNotFoundError):
check_dataset_for_missing_videos(ds)
def test_duplicate_missing_deduped_in_warn_mode(check_dataset_for_missing_videos):
"""Warn mode returns each missing path once."""
ds = _make_video_dataset("/nonexistent/dup.mp4", "/nonexistent/dup.mp4")
with warnings.catch_warnings(record = True):
warnings.simplefilter("always")
missing = check_dataset_for_missing_videos(ds, raise_error = False)
assert missing == ["/nonexistent/dup.mp4"]
# ── Tests: real unsloth_zoo collator integration ─────────────────────────────
# Exercise the real trainer.py subclass against the real zoo base (the fakes
# above don't cover super()/formatting_func); skip when unsloth can't import.
@pytest.fixture(scope = "session")
def real_collator_classes():
try:
from unsloth.trainer import UnslothVisionDataCollator
from unsloth_zoo.vision_utils import (
UnslothVisionDataCollator as ZooBase,
)
except Exception as exc: # noqa: BLE001 - skip on any import failure
pytest.skip(f"full unsloth import unavailable: {exc!r}")
return UnslothVisionDataCollator, ZooBase
def _make_real_collator(real_collator_classes, formatting_func = None):
"""Build the real subclass without its heavy __init__ (needs a processor)."""
subclass, _ = real_collator_classes
collator = subclass.__new__(subclass)
collator.formatting_func = formatting_func
collator._checked_video_paths = set()
return collator
def test_real_collator_blocks_super_on_missing_video(real_collator_classes, monkeypatch):
"""Missing path raises before the base __call__ runs."""
_, zoo_base = real_collator_classes
calls = []
monkeypatch.setattr(zoo_base, "__call__", lambda self, examples: calls.append(examples))
collator = _make_real_collator(real_collator_classes)
with pytest.raises(FileNotFoundError):
collator(_batch("/nonexistent/real.mp4"))
assert calls == [] # base collator was never reached
def test_real_collator_calls_super_with_formatting_disabled(real_collator_classes, monkeypatch):
"""Base must see formatting_func=None and already-formatted examples;
the original formatting_func is restored afterwards."""
seen = {}
def spy(self, examples):
seen["formatting_func"] = self.formatting_func
seen["examples"] = examples
return {"ok": True}
_, zoo_base = real_collator_classes
monkeypatch.setattr(zoo_base, "__call__", spy)
def fmt(example):
return {
"messages": [
{
"role": "user",
"content": [{"type": "video", "video": example["video_id"]}],
}
]
}
with tempfile.NamedTemporaryFile(suffix = ".mp4", delete = False) as f:
f.write(b"x")
tmp = f.name
try:
collator = _make_real_collator(real_collator_classes, formatting_func = fmt)
result = collator([{"video_id": tmp}])
assert result == {"ok": True}
assert seen["formatting_func"] is None
assert seen["examples"][0]["messages"][0]["content"][0]["video"] == tmp
assert collator.formatting_func is fmt
assert tmp in collator._checked_video_paths
finally:
os.unlink(tmp)
def test_real_collator_restores_formatting_func_when_super_raises(
real_collator_classes, monkeypatch
):
"""formatting_func is restored even when the base raises."""
def boom(self, examples):
raise RuntimeError("base collator failed")
_, zoo_base = real_collator_classes
monkeypatch.setattr(zoo_base, "__call__", boom)
def fmt(example):
return {"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]}
collator = _make_real_collator(real_collator_classes, formatting_func = fmt)
with pytest.raises(RuntimeError):
collator([{"anything": 1}])
assert collator.formatting_func is fmt