Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions src/pruna/data/pruna_datamodule.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,9 @@ def from_datasets( # type: ignore[override]
pruna_logger.error("Datasets must contain exactly 3 elements: train, validation, and test.")
raise ValueError()

# Shallow-copy before injecting tokenizer: avoids mutating the caller's dict and
# leaking keys into the shared default from collate_fn_args=dict().
collate_fn_args = dict(collate_fn_args)
if tokenizer is not None:
collate_fn_args["tokenizer"] = tokenizer
if "max_seq_len" not in collate_fn_args:
Expand Down Expand Up @@ -376,6 +379,11 @@ def get_collate_fn(collate_fn_name: str, collate_fn_args: dict) -> Callable:
if missing_required_params:
raise ValueError(f"The following required parameters are missing in collate_fn_args: {missing_required_params}")

# Drop kwargs the collate function does not accept (e.g. tokenizer forwarded for prompt_collate).
has_var_keyword = any(param.kind == inspect.Parameter.VAR_KEYWORD for param in signature.parameters.values())
if not has_var_keyword: # if the signature does not contain **kwargs
collate_fn_args = {k: v for k, v in collate_fn_args.items() if k in signature.parameters}

# Create a partial with the given arguments
collate_fn = partial(collate_fn, **collate_fn_args)
return collate_fn
40 changes: 40 additions & 0 deletions tests/config/test_smashconfig_data.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
from typing import Any, Callable
from unittest.mock import MagicMock
import inspect

import pytest
from datasets import Dataset

from pruna import SmashConfig
from pruna.data.datasets.image import setup_imagenet_dataset
Expand Down Expand Up @@ -61,3 +64,40 @@ def test_img_args_override(
image, _ = next(iter(dataloader))
assert image.shape[2] == args_override["img_size"]
assert image.shape[3] == args_override["img_size"]


@pytest.mark.cpu
def test_add_prompt_data_ignores_unrelated_tokenizer() -> None:
"""Tokenizer on SmashConfig must not break prompt_collate (e.g. after loading a diffusers model)."""
ds = Dataset.from_dict({"text": ["a cat", "a dog"]})
smash_config = SmashConfig()
tokenizer = MagicMock()
tokenizer.model_max_length = 77
smash_config.tokenizer = tokenizer
smash_config.add_data((ds, ds, ds), "prompt_collate")
assert smash_config.data is not None
prompts, none = next(iter(smash_config.data.train_dataloader(batch_size=2, shuffle=False)))
assert sorted(prompts) == ["a cat", "a dog"]
assert none is None


@pytest.mark.cpu
def test_from_datasets_does_not_mutate_collate_fn_args() -> None:
"""from_datasets must copy collate_fn_args before injecting tokenizer."""
ds = Dataset.from_dict({"text": ["a cat", "a dog"]})
datasets = (ds, ds, ds)
tokenizer = MagicMock()
tokenizer.model_max_length = 77

caller_args: dict[str, object] = {}
PrunaDataModule.from_datasets(
datasets,
"text_generation_collate",
tokenizer=tokenizer,
collate_fn_args=caller_args,
)

assert caller_args == {}

default_args = inspect.signature(PrunaDataModule.from_datasets).parameters["collate_fn_args"].default
assert default_args == {}
Loading