Added validation loss

This commit is contained in:
Jaret Burkett
2026-07-19 19:33:45 -06:00
parent cd677c70b5
commit 1eb97b7443
7 changed files with 392 additions and 16 deletions

View File

@@ -337,6 +337,23 @@ class AdapterConfig:
self.i2v_do_start_frame: bool = kwargs.get('i2v_do_start_frame', False)
class ValidationItem:
def __init__(self, **kwargs):
self.image_path: str = kwargs.get('image_path', '')
self.prompt: str = kwargs.get('prompt', '')
class ValidationConfig:
def __init__(self, **kwargs):
self.validation_items: List[ValidationItem] = [
item if isinstance(item, ValidationItem) else ValidationItem(**item)
for item in kwargs.get('validation_items', [])
]
self.resolution: int = kwargs.get('resolution', 512)
self.validate_every_n_steps: int = kwargs.get('validate_every_n_steps', 10)
self.validation_sigmas: List[float] = kwargs.get('validation_sigmas', [1.0, 0.75, 0.5, 0.25])
class EmbeddingConfig:
def __init__(self, **kwargs):
self.trigger = kwargs.get('trigger', 'custom_embedding')
@@ -592,6 +609,10 @@ class TrainConfig:
self.max_loss_debug: bool = kwargs.get("max_loss_debug", False)
# will clip the loss to this amount to prevent wild outliers
self.max_loss: Optional[float] = kwargs.get("max_loss", None)
self.validation_config: Optional[ValidationConfig] = None
validation = kwargs.get('validation_config', None)
if validation is not None:
self.validation_config: ValidationConfig = ValidationConfig(**validation)
ModelArch = Literal['sd1', 'sd2', 'sd3', 'sdxl', 'pixart', 'pixart_sigma', 'auraflow', 'flux', 'flex1', 'flex2', 'lumina2', 'vega', 'ssd', 'wan21', 'anima']
@@ -1417,5 +1438,3 @@ def validate_configs(
if train_config.batch_size > 1 and any(dataset_config.auto_frame_count for dataset_config in dataset_configs):
raise ValueError("Cannot use batch size greater than 1 with auto_frame_count. Please set batch_size to 1 or auto_frame_count to False.")