Bases: DataLoaderArgs
Model for PyTorch data loader arguments.
Source code in src/guidellm/schemas/data/loaders/torch.py
| @DataLoaderArgs.register("pytorch")
class TorchDataLoaderArgs(DataLoaderArgs):
"""Model for PyTorch data loader arguments."""
kind: Literal["pytorch"] = Field( # type: ignore[assignment]
default="pytorch",
description="Type identifier for the generative data loader.",
)
shuffle: bool = Field(
default=False,
description="Shuffle data rows at every epoch.",
)
num_workers: int = Field(
default=1,
description=(
"Number of worker processes for data loading. If 0, data loading "
"will be performed in the main process."
),
)
prefetch_factor: int = Field(
default=4096,
description=(
"Number of samples loaded in advance by each worker. "
"Increasing this generates more data ahead of demand."
),
)
@field_validator("num_workers", mode="after")
@classmethod
def warn_if_changed(cls, v: int) -> int:
if v != 1:
logger.warning(
"The value of data_loader.num_workers has been changed from its "
"default value. This is currently not supported and may lead to "
"unexpected behavior."
)
return v
|