kiln_ai.datamodel.dataset_split
Tools for splitting datasets into train/test/validation splits. Includes filters for selecting which task runs to include in each split.
1""" 2Tools for splitting datasets into train/test/validation splits. Includes filters for selecting which task runs to include in each split. 3""" 4 5import math 6import random 7from typing import TYPE_CHECKING 8 9from pydantic import BaseModel, Field, model_validator 10 11from kiln_ai.datamodel.basemodel import FilenameString, KilnParentedModel 12from kiln_ai.datamodel.dataset_filters import ( 13 DatasetFilter, 14 DatasetFilterId, 15 dataset_filter_from_id, 16) 17from kiln_ai.datamodel.run_config import KilnAgentRunConfigProperties 18from kiln_ai.datamodel.task_run import TaskRun 19 20if TYPE_CHECKING: 21 from kiln_ai.datamodel.task import Task 22 23 24class DatasetToolInfo(BaseModel): 25 """ 26 Information about tools used across task runs in a dataset split. 27 """ 28 29 has_tool_mismatch: bool = Field( 30 description="Whether the tools from each run match across all runs in the dataset split." 31 ) 32 tools: list[str] | None = Field( 33 default_factory=list, 34 description="Common tool IDs shared by every run. Empty means every run has no tools. None means the dataset contains mismatched tool sets.", 35 ) 36 37 38class DatasetSplitDefinition(BaseModel): 39 """ 40 A definition of a split in a dataset. 41 42 Example: name="train", description="The training set", percentage=0.8 (80% of the dataset) 43 """ 44 45 name: FilenameString = Field( 46 description="The name of the dataset split definition." 47 ) 48 description: str | None = Field( 49 default=None, 50 description="A description of the dataset for you and your team. Not used in training.", 51 ) 52 percentage: float = Field( 53 ge=0.0, 54 le=1.0, 55 description="The percentage of the dataset that this split represents (between 0 and 1).", 56 ) 57 58 59AllSplitDefinition: list[DatasetSplitDefinition] = [ 60 DatasetSplitDefinition(name="all", percentage=1.0) 61] 62Train80Test20SplitDefinition: list[DatasetSplitDefinition] = [ 63 DatasetSplitDefinition(name="train", percentage=0.8), 64 DatasetSplitDefinition(name="test", percentage=0.2), 65] 66Train80Val20SplitDefinition: list[DatasetSplitDefinition] = [ 67 DatasetSplitDefinition(name="train", percentage=0.8), 68 DatasetSplitDefinition(name="val", percentage=0.2), 69] 70Train60Test20Val20SplitDefinition: list[DatasetSplitDefinition] = [ 71 DatasetSplitDefinition(name="train", percentage=0.6), 72 DatasetSplitDefinition(name="test", percentage=0.2), 73 DatasetSplitDefinition(name="val", percentage=0.2), 74] 75Train80Test10Val10SplitDefinition: list[DatasetSplitDefinition] = [ 76 DatasetSplitDefinition(name="train", percentage=0.8), 77 DatasetSplitDefinition(name="test", percentage=0.1), 78 DatasetSplitDefinition(name="val", percentage=0.1), 79] 80 81 82class DatasetSplit(KilnParentedModel): 83 """ 84 A collection of task runs, with optional splits (train, test, validation). 85 86 Used to freeze a dataset into train/test/validation splits for repeatable fine-tuning or other tasks. 87 88 Maintains a list of IDs for each split, to avoid data duplication. 89 """ 90 91 name: FilenameString = Field(description="The name of the dataset split.") 92 description: str | None = Field( 93 default=None, 94 description="A description of the dataset for you and your team. Not used in training.", 95 ) 96 splits: list[DatasetSplitDefinition] = Field( 97 default_factory=list, 98 description="The splits in the dataset.", 99 ) 100 split_contents: dict[str, list[str]] = Field( 101 description="The contents of each split in the dataset. The key is the split name, and the value is a list of task run IDs.", 102 ) 103 filter: DatasetFilterId | None = Field( 104 default=None, 105 description="The filter used to build the dataset.", 106 ) 107 108 @model_validator(mode="after") 109 def validate_split_percentages(self) -> "DatasetSplit": 110 total = sum(split.percentage for split in self.splits) 111 if not math.isclose(total, 1.0, rel_tol=1e-9): 112 raise ValueError(f"The sum of split percentages must be 1.0 (got {total})") 113 return self 114 115 @classmethod 116 def from_task( 117 cls, 118 name: str, 119 task: "Task", 120 splits: list[DatasetSplitDefinition], 121 filter_id: DatasetFilterId = "all", 122 description: str | None = None, 123 ): 124 """ 125 Build a dataset split from a task. 126 """ 127 filter = dataset_filter_from_id(filter_id) 128 split_contents = cls.build_split_contents(task, splits, filter) 129 return cls( 130 parent=task, 131 name=name, 132 description=description, 133 splits=splits, 134 split_contents=split_contents, 135 filter=filter_id, 136 ) 137 138 @classmethod 139 def build_split_contents( 140 cls, 141 task: "Task", 142 splits: list[DatasetSplitDefinition], 143 filter: DatasetFilter, 144 ) -> dict[str, list[str]]: 145 # Datasets must not include intermediate multiturn turns — only leaves. 146 runs = list(task.runs()) 147 valid_ids = [] 148 for task_run in runs: 149 if filter(task_run): 150 valid_ids.append(task_run.id) 151 152 # Shuffle and split by split percentage 153 random.shuffle(valid_ids) 154 split_contents = {} 155 start_idx = 0 156 remaining_items = len(valid_ids) 157 158 # Handle all splits except the last one 159 for split in splits[:-1]: 160 split_size = round(len(valid_ids) * split.percentage) 161 split_contents[split.name] = valid_ids[start_idx : start_idx + split_size] 162 start_idx += split_size 163 remaining_items -= split_size 164 165 # Last split gets all remaining items (for rounding) 166 if splits: 167 split_contents[splits[-1].name] = valid_ids[start_idx:] 168 169 return split_contents 170 171 def parent_task(self) -> "Task | None": 172 # inline import to avoid circular import 173 from kiln_ai.datamodel import Task 174 175 if not isinstance(self.parent, Task): 176 return None 177 return self.parent 178 179 def missing_count(self) -> int: 180 """ 181 Returns: 182 int: the number of task runs that have an ID persisted in this dataset split, but no longer exist in the dataset 183 """ 184 parent = self.parent_task() 185 if parent is None: 186 raise ValueError("DatasetSplit has no parent task") 187 188 # Include intermediate runs: a split snapshots run IDs at creation 189 # time (leaves-only), but the user can later extend a multiturn 190 # conversation - turning a run that was a leaf when the split was 191 # built into an intermediate. The run is still on disk, so it 192 # shouldn't be reported as "missing". 193 runs = parent.runs(include_intermediate_runs=True, readonly=True) 194 all_ids = set(run.id for run in runs) 195 all_ids_in_splits = set() 196 for ids in self.split_contents.values(): 197 all_ids_in_splits.update(ids) 198 missing = all_ids_in_splits - all_ids 199 return len(missing) 200 201 def _get_runs(self) -> list[TaskRun]: 202 """ 203 Get all task runs referenced in this dataset split. 204 205 Returns: 206 list[TaskRun]: list of task runs in this dataset split 207 """ 208 parent = self.parent_task() 209 if parent is None: 210 return [] 211 212 runs = [] 213 all_run_ids = set() 214 for run_ids in self.split_contents.values(): 215 all_run_ids.update(run_ids) 216 217 # Include intermediate runs: a split snapshots run IDs at creation 218 # time (leaves-only), but the user can later extend a multiturn 219 # conversation - turning a run that was a leaf when the split was 220 # built into an intermediate. We still want to resolve it. 221 for task_run in parent.runs(include_intermediate_runs=True, readonly=True): 222 if task_run.id in all_run_ids: 223 runs.append(task_run) 224 225 return runs 226 227 @staticmethod 228 def compute_tool_info(runs: list[TaskRun]) -> DatasetToolInfo: 229 """ 230 Compute tool info from a list of task runs. 231 232 Args: 233 runs: list of task runs to analyze 234 235 Returns: 236 DatasetToolInfo: information about tools used across the task runs 237 """ 238 239 has_tool_mismatch = False 240 tools: set[str] | None = None 241 242 for run in runs: 243 # Extract tools from run config, treating missing source/run_config/tools_config as empty tools 244 run_tools: set[str] = set() 245 source = run.output.source if run.output else None 246 if source is not None and isinstance( 247 source.run_config, KilnAgentRunConfigProperties 248 ): 249 tools_config = source.run_config.tools_config 250 if tools_config is not None: 251 run_tools = set(tools_config.tools) 252 253 # First run establishes the expected tool set (including empty) 254 if tools is None: 255 tools = run_tools 256 elif run_tools != tools: 257 # Mismatch found 258 has_tool_mismatch = True 259 tools = None 260 break 261 262 # If no valid runs were processed, return empty tools 263 if tools is None: 264 if not has_tool_mismatch: 265 tools = set() 266 267 return DatasetToolInfo( 268 has_tool_mismatch=has_tool_mismatch, 269 tools=None if tools is None else sorted(tools), 270 ) 271 272 def tool_info(self) -> DatasetToolInfo: 273 """ 274 Helper method to compute tool info for the dataset split. Iterate through all runs in the dataset split and check the tools used in each run config. 275 276 Returns: 277 DatasetToolInfo: information about tools used across task runs in this dataset split 278 """ 279 runs = self._get_runs() 280 tool_info = self.compute_tool_info(runs) 281 return tool_info
25class DatasetToolInfo(BaseModel): 26 """ 27 Information about tools used across task runs in a dataset split. 28 """ 29 30 has_tool_mismatch: bool = Field( 31 description="Whether the tools from each run match across all runs in the dataset split." 32 ) 33 tools: list[str] | None = Field( 34 default_factory=list, 35 description="Common tool IDs shared by every run. Empty means every run has no tools. None means the dataset contains mismatched tool sets.", 36 )
Information about tools used across task runs in a dataset split.
39class DatasetSplitDefinition(BaseModel): 40 """ 41 A definition of a split in a dataset. 42 43 Example: name="train", description="The training set", percentage=0.8 (80% of the dataset) 44 """ 45 46 name: FilenameString = Field( 47 description="The name of the dataset split definition." 48 ) 49 description: str | None = Field( 50 default=None, 51 description="A description of the dataset for you and your team. Not used in training.", 52 ) 53 percentage: float = Field( 54 ge=0.0, 55 le=1.0, 56 description="The percentage of the dataset that this split represents (between 0 and 1).", 57 )
A definition of a split in a dataset.
Example: name="train", description="The training set", percentage=0.8 (80% of the dataset)
83class DatasetSplit(KilnParentedModel): 84 """ 85 A collection of task runs, with optional splits (train, test, validation). 86 87 Used to freeze a dataset into train/test/validation splits for repeatable fine-tuning or other tasks. 88 89 Maintains a list of IDs for each split, to avoid data duplication. 90 """ 91 92 name: FilenameString = Field(description="The name of the dataset split.") 93 description: str | None = Field( 94 default=None, 95 description="A description of the dataset for you and your team. Not used in training.", 96 ) 97 splits: list[DatasetSplitDefinition] = Field( 98 default_factory=list, 99 description="The splits in the dataset.", 100 ) 101 split_contents: dict[str, list[str]] = Field( 102 description="The contents of each split in the dataset. The key is the split name, and the value is a list of task run IDs.", 103 ) 104 filter: DatasetFilterId | None = Field( 105 default=None, 106 description="The filter used to build the dataset.", 107 ) 108 109 @model_validator(mode="after") 110 def validate_split_percentages(self) -> "DatasetSplit": 111 total = sum(split.percentage for split in self.splits) 112 if not math.isclose(total, 1.0, rel_tol=1e-9): 113 raise ValueError(f"The sum of split percentages must be 1.0 (got {total})") 114 return self 115 116 @classmethod 117 def from_task( 118 cls, 119 name: str, 120 task: "Task", 121 splits: list[DatasetSplitDefinition], 122 filter_id: DatasetFilterId = "all", 123 description: str | None = None, 124 ): 125 """ 126 Build a dataset split from a task. 127 """ 128 filter = dataset_filter_from_id(filter_id) 129 split_contents = cls.build_split_contents(task, splits, filter) 130 return cls( 131 parent=task, 132 name=name, 133 description=description, 134 splits=splits, 135 split_contents=split_contents, 136 filter=filter_id, 137 ) 138 139 @classmethod 140 def build_split_contents( 141 cls, 142 task: "Task", 143 splits: list[DatasetSplitDefinition], 144 filter: DatasetFilter, 145 ) -> dict[str, list[str]]: 146 # Datasets must not include intermediate multiturn turns — only leaves. 147 runs = list(task.runs()) 148 valid_ids = [] 149 for task_run in runs: 150 if filter(task_run): 151 valid_ids.append(task_run.id) 152 153 # Shuffle and split by split percentage 154 random.shuffle(valid_ids) 155 split_contents = {} 156 start_idx = 0 157 remaining_items = len(valid_ids) 158 159 # Handle all splits except the last one 160 for split in splits[:-1]: 161 split_size = round(len(valid_ids) * split.percentage) 162 split_contents[split.name] = valid_ids[start_idx : start_idx + split_size] 163 start_idx += split_size 164 remaining_items -= split_size 165 166 # Last split gets all remaining items (for rounding) 167 if splits: 168 split_contents[splits[-1].name] = valid_ids[start_idx:] 169 170 return split_contents 171 172 def parent_task(self) -> "Task | None": 173 # inline import to avoid circular import 174 from kiln_ai.datamodel import Task 175 176 if not isinstance(self.parent, Task): 177 return None 178 return self.parent 179 180 def missing_count(self) -> int: 181 """ 182 Returns: 183 int: the number of task runs that have an ID persisted in this dataset split, but no longer exist in the dataset 184 """ 185 parent = self.parent_task() 186 if parent is None: 187 raise ValueError("DatasetSplit has no parent task") 188 189 # Include intermediate runs: a split snapshots run IDs at creation 190 # time (leaves-only), but the user can later extend a multiturn 191 # conversation - turning a run that was a leaf when the split was 192 # built into an intermediate. The run is still on disk, so it 193 # shouldn't be reported as "missing". 194 runs = parent.runs(include_intermediate_runs=True, readonly=True) 195 all_ids = set(run.id for run in runs) 196 all_ids_in_splits = set() 197 for ids in self.split_contents.values(): 198 all_ids_in_splits.update(ids) 199 missing = all_ids_in_splits - all_ids 200 return len(missing) 201 202 def _get_runs(self) -> list[TaskRun]: 203 """ 204 Get all task runs referenced in this dataset split. 205 206 Returns: 207 list[TaskRun]: list of task runs in this dataset split 208 """ 209 parent = self.parent_task() 210 if parent is None: 211 return [] 212 213 runs = [] 214 all_run_ids = set() 215 for run_ids in self.split_contents.values(): 216 all_run_ids.update(run_ids) 217 218 # Include intermediate runs: a split snapshots run IDs at creation 219 # time (leaves-only), but the user can later extend a multiturn 220 # conversation - turning a run that was a leaf when the split was 221 # built into an intermediate. We still want to resolve it. 222 for task_run in parent.runs(include_intermediate_runs=True, readonly=True): 223 if task_run.id in all_run_ids: 224 runs.append(task_run) 225 226 return runs 227 228 @staticmethod 229 def compute_tool_info(runs: list[TaskRun]) -> DatasetToolInfo: 230 """ 231 Compute tool info from a list of task runs. 232 233 Args: 234 runs: list of task runs to analyze 235 236 Returns: 237 DatasetToolInfo: information about tools used across the task runs 238 """ 239 240 has_tool_mismatch = False 241 tools: set[str] | None = None 242 243 for run in runs: 244 # Extract tools from run config, treating missing source/run_config/tools_config as empty tools 245 run_tools: set[str] = set() 246 source = run.output.source if run.output else None 247 if source is not None and isinstance( 248 source.run_config, KilnAgentRunConfigProperties 249 ): 250 tools_config = source.run_config.tools_config 251 if tools_config is not None: 252 run_tools = set(tools_config.tools) 253 254 # First run establishes the expected tool set (including empty) 255 if tools is None: 256 tools = run_tools 257 elif run_tools != tools: 258 # Mismatch found 259 has_tool_mismatch = True 260 tools = None 261 break 262 263 # If no valid runs were processed, return empty tools 264 if tools is None: 265 if not has_tool_mismatch: 266 tools = set() 267 268 return DatasetToolInfo( 269 has_tool_mismatch=has_tool_mismatch, 270 tools=None if tools is None else sorted(tools), 271 ) 272 273 def tool_info(self) -> DatasetToolInfo: 274 """ 275 Helper method to compute tool info for the dataset split. Iterate through all runs in the dataset split and check the tools used in each run config. 276 277 Returns: 278 DatasetToolInfo: information about tools used across task runs in this dataset split 279 """ 280 runs = self._get_runs() 281 tool_info = self.compute_tool_info(runs) 282 return tool_info
A collection of task runs, with optional splits (train, test, validation).
Used to freeze a dataset into train/test/validation splits for repeatable fine-tuning or other tasks.
Maintains a list of IDs for each split, to avoid data duplication.
109 @model_validator(mode="after") 110 def validate_split_percentages(self) -> "DatasetSplit": 111 total = sum(split.percentage for split in self.splits) 112 if not math.isclose(total, 1.0, rel_tol=1e-9): 113 raise ValueError(f"The sum of split percentages must be 1.0 (got {total})") 114 return self
116 @classmethod 117 def from_task( 118 cls, 119 name: str, 120 task: "Task", 121 splits: list[DatasetSplitDefinition], 122 filter_id: DatasetFilterId = "all", 123 description: str | None = None, 124 ): 125 """ 126 Build a dataset split from a task. 127 """ 128 filter = dataset_filter_from_id(filter_id) 129 split_contents = cls.build_split_contents(task, splits, filter) 130 return cls( 131 parent=task, 132 name=name, 133 description=description, 134 splits=splits, 135 split_contents=split_contents, 136 filter=filter_id, 137 )
Build a dataset split from a task.
139 @classmethod 140 def build_split_contents( 141 cls, 142 task: "Task", 143 splits: list[DatasetSplitDefinition], 144 filter: DatasetFilter, 145 ) -> dict[str, list[str]]: 146 # Datasets must not include intermediate multiturn turns — only leaves. 147 runs = list(task.runs()) 148 valid_ids = [] 149 for task_run in runs: 150 if filter(task_run): 151 valid_ids.append(task_run.id) 152 153 # Shuffle and split by split percentage 154 random.shuffle(valid_ids) 155 split_contents = {} 156 start_idx = 0 157 remaining_items = len(valid_ids) 158 159 # Handle all splits except the last one 160 for split in splits[:-1]: 161 split_size = round(len(valid_ids) * split.percentage) 162 split_contents[split.name] = valid_ids[start_idx : start_idx + split_size] 163 start_idx += split_size 164 remaining_items -= split_size 165 166 # Last split gets all remaining items (for rounding) 167 if splits: 168 split_contents[splits[-1].name] = valid_ids[start_idx:] 169 170 return split_contents
180 def missing_count(self) -> int: 181 """ 182 Returns: 183 int: the number of task runs that have an ID persisted in this dataset split, but no longer exist in the dataset 184 """ 185 parent = self.parent_task() 186 if parent is None: 187 raise ValueError("DatasetSplit has no parent task") 188 189 # Include intermediate runs: a split snapshots run IDs at creation 190 # time (leaves-only), but the user can later extend a multiturn 191 # conversation - turning a run that was a leaf when the split was 192 # built into an intermediate. The run is still on disk, so it 193 # shouldn't be reported as "missing". 194 runs = parent.runs(include_intermediate_runs=True, readonly=True) 195 all_ids = set(run.id for run in runs) 196 all_ids_in_splits = set() 197 for ids in self.split_contents.values(): 198 all_ids_in_splits.update(ids) 199 missing = all_ids_in_splits - all_ids 200 return len(missing)
Returns: int: the number of task runs that have an ID persisted in this dataset split, but no longer exist in the dataset
228 @staticmethod 229 def compute_tool_info(runs: list[TaskRun]) -> DatasetToolInfo: 230 """ 231 Compute tool info from a list of task runs. 232 233 Args: 234 runs: list of task runs to analyze 235 236 Returns: 237 DatasetToolInfo: information about tools used across the task runs 238 """ 239 240 has_tool_mismatch = False 241 tools: set[str] | None = None 242 243 for run in runs: 244 # Extract tools from run config, treating missing source/run_config/tools_config as empty tools 245 run_tools: set[str] = set() 246 source = run.output.source if run.output else None 247 if source is not None and isinstance( 248 source.run_config, KilnAgentRunConfigProperties 249 ): 250 tools_config = source.run_config.tools_config 251 if tools_config is not None: 252 run_tools = set(tools_config.tools) 253 254 # First run establishes the expected tool set (including empty) 255 if tools is None: 256 tools = run_tools 257 elif run_tools != tools: 258 # Mismatch found 259 has_tool_mismatch = True 260 tools = None 261 break 262 263 # If no valid runs were processed, return empty tools 264 if tools is None: 265 if not has_tool_mismatch: 266 tools = set() 267 268 return DatasetToolInfo( 269 has_tool_mismatch=has_tool_mismatch, 270 tools=None if tools is None else sorted(tools), 271 )
Compute tool info from a list of task runs.
Args: runs: list of task runs to analyze
Returns: DatasetToolInfo: information about tools used across the task runs
273 def tool_info(self) -> DatasetToolInfo: 274 """ 275 Helper method to compute tool info for the dataset split. Iterate through all runs in the dataset split and check the tools used in each run config. 276 277 Returns: 278 DatasetToolInfo: information about tools used across task runs in this dataset split 279 """ 280 runs = self._get_runs() 281 tool_info = self.compute_tool_info(runs) 282 return tool_info
Helper method to compute tool info for the dataset split. Iterate through all runs in the dataset split and check the tools used in each run config.
Returns: DatasetToolInfo: information about tools used across task runs in this dataset split
The type of the None singleton.
Configuration for the model, should be a dictionary conforming to [ConfigDict][pydantic.config.ConfigDict].
365def init_private_attributes(self: BaseModel, context: Any, /) -> None: 366 """This function is meant to behave like a BaseModel method to initialize private attributes. 367 368 It takes context as an argument since that's what pydantic-core passes when calling it. 369 370 Args: 371 self: The BaseModel instance. 372 context: The context. 373 """ 374 if getattr(self, '__pydantic_private__', None) is None: 375 pydantic_private = {} 376 for name, private_attr in self.__private_attributes__.items(): 377 # Avoid needlessly creating a new dict for the validated data: 378 if private_attr.default_factory_takes_validated_data: 379 default = private_attr.get_default( 380 call_default_factory=True, validated_data={**self.__dict__, **pydantic_private} 381 ) 382 else: 383 default = private_attr.get_default(call_default_factory=True) 384 if default is not PydanticUndefined: 385 pydantic_private[name] = default 386 object_setattr(self, '__pydantic_private__', pydantic_private)
This function is meant to behave like a BaseModel method to initialize private attributes.
It takes context as an argument since that's what pydantic-core passes when calling it.
Args: self: The BaseModel instance. context: The context.