Coverage for packages/dqm-ml-core/src/dqm_ml_core/models/dataloaders.py: 100%

46 statements  

« prev     ^ index     » next       coverage.py v7.14.1, created at 2026-07-21 08:27 +0000

1"""Data loading and transformation configuration models. 

2 

3Defines models for filters, path sampling, data splitting, column transformations, 

4and dataloader configurations for Parquet and CSV sources. 

5""" 

6 

7from enum import Enum 

8from typing import Literal 

9 

10from pydantic import BaseModel, ConfigDict, Field 

11 

12from dqm_ml_core.models.global_ import StorageConfig 

13 

14 

15class FilterConfig(BaseModel): 

16 """Row-level filter applied during data loading.""" 

17 

18 model_config = ConfigDict(extra="forbid") 

19 

20 column: str = Field(description="Column name to filter on.") 

21 values: list[bool] | list[str] | list[int] | list[float] = Field( 

22 description="Value(s) to keep. Rows where column matches are included.", 

23 ) 

24 

25 

26class SamplePathConfig(BaseModel): 

27 """Per-column path prefix configuration.""" 

28 

29 model_config = ConfigDict(extra="forbid") 

30 

31 column: str = Field(description="Column name containing relative file paths.") 

32 prefix: str | None = Field(default=None, description="Base directory for resolving relative paths.") 

33 

34 

35class SplitConfig(BaseModel): 

36 """How to split data into named groups (e.g. train / test).""" 

37 

38 model_config = ConfigDict(extra="forbid") 

39 

40 by: str = Field(description="Column used to determine the split group.") 

41 values: list[str] | None = Field( 

42 default=None, 

43 description="Explicit list of split-group values to materialise. Auto-discovered if None.", 

44 ) 

45 exclude: list[str] | None = Field( 

46 default=None, 

47 description="Split-group values to exclude (fnmatch patterns supported).", 

48 ) 

49 

50 

51class TransformType(str, Enum): 

52 """Target data type for column transformations.""" 

53 

54 INT32 = "int32" 

55 INT64 = "int64" 

56 FLOAT32 = "float32" 

57 FLOAT64 = "float64" 

58 BOOL = "bool" 

59 STR = "str" 

60 CATEGORICAL = "categorical" 

61 

62 

63class TransformConfig(BaseModel): 

64 """Column type-casting transformation.""" 

65 

66 model_config = ConfigDict(extra="forbid") 

67 

68 column: str = Field(description="Column name to transform.") 

69 to_type: TransformType = Field(description="Target data type.") 

70 in_place: bool = Field(default=False, description="Overwrite the original column in place.") 

71 

72 

73class DataLoaderConfig(BaseModel): 

74 """Configuration for a single dataloader (Parquet or CSV).""" 

75 

76 model_config = ConfigDict(extra="forbid") 

77 

78 name: str = Field(description="Unique name for this dataloader.") 

79 type: Literal["parquet", "csv"] = Field(description="Data file format.") 

80 path: str = Field(description="Glob pattern or path to data files.") 

81 id_column: str | None = Field(default=None, description="Column used as row identifier.") 

82 batch_size: int = Field(default=10000, description="Number of rows per batch.") 

83 filters: list[FilterConfig] | None = None 

84 sample_path: list[SamplePathConfig] | None = None 

85 split: SplitConfig | None = None 

86 transform: list[TransformConfig] | None = None 

87 storage: StorageConfig | None = None 

88 

89 

90class DataLoadersConfig(BaseModel): 

91 """Collection of dataloaders for a job.""" 

92 

93 model_config = ConfigDict(extra="forbid") 

94 

95 storage: StorageConfig | None = Field( 

96 default=None, 

97 description="Default storage config inherited by all loaders.", 

98 ) 

99 loaders: list[DataLoaderConfig] = Field(description="List of dataloader configurations.")