Skip to content

Commit 823833f

Browse files
author
bghira
committed
Fix validation artifact asset formats
1 parent c8ebd18 commit 823833f

4 files changed

Lines changed: 185 additions & 47 deletions

File tree

simpletuner/helpers/publishing/huggingface.py

Lines changed: 70 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
from simpletuner.helpers.publishing.metadata import save_model_card, save_training_config
1010
from simpletuner.helpers.training.state_tracker import StateTracker
11+
from simpletuner.helpers.training.validation_images import save_image_with_validation_format
1112

1213
logger = logging.getLogger(__name__)
1314
from simpletuner.helpers.training.multi_process import should_log
@@ -21,6 +22,10 @@
2122
LORA_SAFETENSORS_FILENAME = "pytorch_lora_weights.safetensors"
2223
EMA_SAFETENSORS_FILENAME = "ema_model.safetensors"
2324
SLA_ATTENTION_FILENAME = "sla_attention.pt"
25+
VALIDATION_IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp", ".gif"}
26+
VALIDATION_VIDEO_SUFFIXES = {".mp4", ".avi", ".mov", ".webm"}
27+
VALIDATION_AUDIO_SUFFIXES = {".wav", ".flac", ".mp3", ".ogg", ".m4a"}
28+
VALIDATION_MEDIA_SUFFIXES = VALIDATION_IMAGE_SUFFIXES | VALIDATION_VIDEO_SUFFIXES | VALIDATION_AUDIO_SUFFIXES
2429

2530

2631
class HubManager:
@@ -409,6 +414,59 @@ def find_latest_checkpoint(self):
409414

410415
return highest_checkpoint
411416

417+
@staticmethod
418+
def _validation_media_shortname(filename: str, checkpoint_step: int) -> str | None:
419+
stem = Path(filename).stem
420+
prefix = f"step_{checkpoint_step}_"
421+
if not stem.startswith(prefix):
422+
return None
423+
424+
remainder = stem[len(prefix) :]
425+
image_or_video_parts = remainder.rsplit("_", 2)
426+
if len(image_or_video_parts) == 3 and image_or_video_parts[1].isdigit():
427+
return image_or_video_parts[0]
428+
429+
audio_parts = remainder.rsplit("_", 1)
430+
if len(audio_parts) == 2 and audio_parts[1].isdigit():
431+
return audio_parts[0]
432+
433+
return None
434+
435+
def _filter_checkpoint_validation_media(
436+
self,
437+
checkpoint_step: int,
438+
validation_images: dict | None,
439+
validation_audios: dict | None,
440+
) -> tuple[dict, dict]:
441+
filtered_images = {}
442+
filtered_audios = {}
443+
validation_dir = Path(self.config.output_dir) / "validation_images"
444+
if not validation_dir.exists():
445+
return filtered_images, filtered_audios
446+
447+
image_shortnames = set(validation_images.keys()) if validation_images else set()
448+
audio_shortnames = set(validation_audios.keys()) if validation_audios else set()
449+
shortnames = image_shortnames | audio_shortnames
450+
if not shortnames:
451+
return filtered_images, filtered_audios
452+
453+
for media_path in sorted(validation_dir.iterdir(), key=lambda path: path.name):
454+
if not media_path.is_file() or media_path.suffix.lower() not in VALIDATION_MEDIA_SUFFIXES:
455+
continue
456+
shortname = self._validation_media_shortname(media_path.name, checkpoint_step)
457+
if shortname not in shortnames:
458+
continue
459+
460+
suffix = media_path.suffix.lower()
461+
if suffix in VALIDATION_AUDIO_SUFFIXES:
462+
if shortname in audio_shortnames:
463+
filtered_audios.setdefault(shortname, []).append(str(media_path))
464+
continue
465+
if shortname in image_shortnames:
466+
filtered_images.setdefault(shortname, []).append(str(media_path))
467+
468+
return filtered_images, filtered_audios
469+
412470
def upload_latest_checkpoint(
413471
self,
414472
validation_images: dict,
@@ -435,30 +493,11 @@ def upload_latest_checkpoint(
435493
filtered_images = {}
436494
filtered_audios = {}
437495
if (validation_images or validation_audios) and checkpoint_step is not None:
438-
validation_dir = os.path.join(self.config.output_dir, "validation_images")
439-
if os.path.exists(validation_dir):
440-
shortnames = set()
441-
if validation_images:
442-
shortnames.update(validation_images.keys())
443-
if validation_audios:
444-
shortnames.update(validation_audios.keys())
445-
for shortname in shortnames:
446-
step_pattern = f"step_{checkpoint_step}_"
447-
for img_file in os.listdir(validation_dir):
448-
if step_pattern in img_file and shortname in img_file:
449-
img_path = os.path.join(validation_dir, img_file)
450-
try:
451-
if img_path.endswith((".mp4", ".avi", ".mov", ".webm")):
452-
filtered_images.setdefault(shortname, []).append(img_path)
453-
elif img_path.endswith((".wav", ".flac", ".mp3", ".ogg", ".m4a")):
454-
filtered_audios.setdefault(shortname, []).append(img_path)
455-
else:
456-
from PIL import Image
457-
458-
img = Image.open(img_path)
459-
filtered_images.setdefault(shortname, []).append(img)
460-
except Exception as e:
461-
logger.warning(f"Could not load validation asset {img_path}: {e}")
496+
filtered_images, filtered_audios = self._filter_checkpoint_validation_media(
497+
checkpoint_step,
498+
validation_images,
499+
validation_audios,
500+
)
462501

463502
# Only use media generated at this checkpoint step. Previous-step benchmark samples
464503
# should not be associated with the checkpoint model card.
@@ -503,12 +542,13 @@ def upload_validation_images(self, validation_images, webhook_handler=None, over
503542
images = [images]
504543
sub_idx = 0
505544
for image in images:
506-
image_path = os.path.join(
507-
override_path or self.config.output_dir,
508-
"assets",
509-
f"image_{idx}_{sub_idx}.png",
545+
assets_dir = os.path.join(override_path or self.config.output_dir, "assets")
546+
os.makedirs(assets_dir, exist_ok=True)
547+
image_path, image_extension = save_image_with_validation_format(
548+
image,
549+
os.path.join(assets_dir, f"image_{idx}_{sub_idx}"),
550+
self.config,
510551
)
511-
image.save(image_path, format="PNG")
512552
if not self.config.push_to_hub:
513553
continue
514554
attempt = 0
@@ -517,7 +557,7 @@ def upload_validation_images(self, validation_images, webhook_handler=None, over
517557
try:
518558
self._hub_api.upload_file(
519559
repo_id=self._repo_id,
520-
path_in_repo=f"/assets/image_{idx}_{sub_idx}.png",
560+
path_in_repo=f"/assets/image_{idx}_{sub_idx}.{image_extension}",
521561
path_or_fileobj=image_path,
522562
commit_message="Validation image auto-generated by SimpleTuner",
523563
token=self.hub_token,

simpletuner/helpers/publishing/metadata.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
from simpletuner.helpers.data_backend.dataset_types import DatasetType, ensure_dataset_type
1414
from simpletuner.helpers.models.common import AudioModelFoundation, ModelFoundation
1515
from simpletuner.helpers.training.state_tracker import StateTracker
16+
from simpletuner.helpers.training.validation_images import save_image_with_validation_format
1617

1718
logger = logging.getLogger(__name__)
1819
from simpletuner.helpers.training.multi_process import should_log
@@ -488,6 +489,7 @@ def save_metadata_sample(
488489
image_path: str,
489490
image: Union[Image.Image, np.ndarray, list, str, torch.Tensor],
490491
sample_rate: int = 44100,
492+
config: Any = None,
491493
):
492494
if isinstance(image, str):
493495
import shutil
@@ -519,9 +521,7 @@ def save_metadata_sample(
519521
fps=StateTracker.get_args().framerate,
520522
)
521523
elif isinstance(image, Image.Image):
522-
file_extension = "png"
523-
output_path = f"{image_path}.{file_extension}"
524-
image.save(output_path, format="PNG")
524+
output_path, file_extension = save_image_with_validation_format(image, image_path, config)
525525
else:
526526
raise ValueError(f"Cannot export sample type {type(image)} yet.")
527527

@@ -772,6 +772,7 @@ def _add_widget_entries(media, asset_prefix: str):
772772
image_path=os.path.join(assets_folder, f"{asset_prefix}_{idx}_{sub_idx}"),
773773
image=media_sample,
774774
sample_rate=audio_sample_rate,
775+
config=args,
775776
)
776777
asset_filename = os.path.basename(output_path)
777778
if media_extension in {"mp4", "avi", "mov", "webm"}:

simpletuner/helpers/training/validation_images.py

Lines changed: 28 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -2,40 +2,54 @@
22
import os
33

44
import numpy as np
5-
import wandb
65
from PIL import Image
76

7+
import wandb
88
from simpletuner.helpers.training.local_metrics import record_validation_media
99
from simpletuner.helpers.training.state_tracker import StateTracker
1010

1111
logger = logging.getLogger(__name__)
1212

13+
VALIDATION_IMAGE_FORMATS = {"png", "webp", "jpeg"}
1314

14-
def save_validation_image(
15-
image,
16-
save_dir,
17-
filename_stem,
18-
config,
19-
*,
20-
label,
21-
index,
22-
resolution,
23-
):
15+
16+
def validation_image_file_extension(config) -> str:
2417
image_format = str(getattr(config, "validation_image_format", "png") or "png").lower()
25-
if image_format not in {"png", "webp", "jpeg"}:
18+
if image_format not in VALIDATION_IMAGE_FORMATS:
19+
raise ValueError("validation_image_format must be one of: png, webp, jpeg.")
20+
return "jpg" if image_format == "jpeg" else image_format
21+
22+
23+
def save_image_with_validation_format(image, image_path_stem, config):
24+
image_format = str(getattr(config, "validation_image_format", "png") or "png").lower()
25+
if image_format not in VALIDATION_IMAGE_FORMATS:
2626
raise ValueError("validation_image_format must be one of: png, webp, jpeg.")
2727
configured_quality = getattr(config, "validation_image_quality", 90)
2828
quality = 90 if configured_quality is None else int(configured_quality)
2929
if not 1 <= quality <= 100:
3030
raise ValueError("validation_image_quality must be within [1, 100].")
3131

32-
extension = "jpg" if image_format == "jpeg" else image_format
33-
save_path = os.path.join(save_dir, f"{filename_stem}.{extension}")
32+
extension = validation_image_file_extension(config)
33+
save_path = f"{image_path_stem}.{extension}"
3434
save_image = image.convert("RGB") if image_format == "jpeg" and getattr(image, "mode", None) != "RGB" else image
3535
save_kwargs = {"format": image_format.upper()}
3636
if image_format in {"webp", "jpeg"}:
3737
save_kwargs["quality"] = quality
3838
save_image.save(save_path, **save_kwargs)
39+
return save_path, extension
40+
41+
42+
def save_validation_image(
43+
image,
44+
save_dir,
45+
filename_stem,
46+
config,
47+
*,
48+
label,
49+
index,
50+
resolution,
51+
):
52+
save_path, _extension = save_image_with_validation_format(image, os.path.join(save_dir, filename_stem), config)
3953
record_validation_media(
4054
config,
4155
save_path,

tests/test_model_card.py

Lines changed: 83 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from unittest.mock import MagicMock, patch
88

99
import torch
10+
from PIL import Image
1011

1112
from simpletuner.helpers.publishing.huggingface import HubManager
1213
from simpletuner.helpers.publishing.metadata import *
@@ -72,6 +73,8 @@ def setUp(self):
7273
self.args.validation_guidance_skip_layers = None
7374
self.args.validation_seed = 1234
7475
self.args.validation_noise_scheduler = "ddim"
76+
self.args.validation_image_format = "png"
77+
self.args.validation_image_quality = 90
7578
self.args.model_card_safe_for_work = True
7679
self.args.learning_rate = 1e-4
7780
self.args.max_grad_norm = 1.0
@@ -83,6 +86,7 @@ def setUp(self):
8386
self.args.base_model_precision = "no_change"
8487
self.args.flux_guidance_mode = "constant"
8588
self.args.flux_guidance_value = 1.0
89+
self.args.peft_lora_mode = "standard"
8690
self.args.t5_padding = "unmodified"
8791
self.args.enable_xformers_memory_efficient_attention = False
8892
self.args.attention_mechanism = "diffusers"
@@ -561,6 +565,85 @@ def test_upload_full_model_uses_empty_repo_path_for_top_level_uploads(self):
561565
hub_manager._hub_api.upload_folder.assert_called_once()
562566
self.assertEqual(hub_manager._hub_api.upload_folder.call_args.kwargs["path_in_repo"], "")
563567

568+
def test_save_model_card_honors_validation_image_format_for_assets(self):
569+
self.args.model_family = "sdxl"
570+
self.args.model_type = "lora"
571+
self.args.lora_type = "standard"
572+
self.args.validation_image_format = "jpeg"
573+
self.args.validation_image_quality = 81
574+
self.args.controlnet = False
575+
self.args.control = False
576+
model = MagicMock(
577+
MODEL_LICENSE="other",
578+
PREDICTION_TYPE=SimpleNamespace(value="epsilon"),
579+
MODEL_TYPE=SimpleNamespace(value="unet"),
580+
gligen=False,
581+
)
582+
model.validation_audio_sample_rate.return_value = 44100
583+
model.custom_model_card_schedule_info.return_value = ""
584+
585+
with tempfile.TemporaryDirectory() as tmpdir:
586+
with (
587+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_model_family", return_value="sdxl"),
588+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_data_backends", return_value={}),
589+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_weight_dtype", return_value=torch.bfloat16),
590+
patch(
591+
"simpletuner.helpers.publishing.metadata.StateTracker.get_accelerator",
592+
return_value=MagicMock(num_processes=1),
593+
),
594+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_args", return_value=self.args),
595+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_model", return_value=model),
596+
):
597+
save_model_card(
598+
repo_id="test-repo",
599+
images={"prompt": [Image.new("RGBA", (8, 8), color=(255, 0, 0, 128))]},
600+
base_model="test-base-model",
601+
train_text_encoder=False,
602+
prompt="Test prompt",
603+
validation_prompts=["Test prompt"],
604+
validation_shortnames=["prompt"],
605+
repo_folder=tmpdir,
606+
model=model,
607+
global_step=1000,
608+
epoch=1,
609+
)
610+
611+
readme = Path(tmpdir, "README.md").read_text(encoding="utf-8")
612+
self.assertIn("url: ./assets/image_0_0.jpg", readme)
613+
self.assertTrue(Path(tmpdir, "assets", "image_0_0.jpg").exists())
614+
self.assertFalse(Path(tmpdir, "assets", "image_0_0.png").exists())
615+
616+
def test_checkpoint_validation_media_filter_matches_shortnames_exactly(self):
617+
with tempfile.TemporaryDirectory() as tmpdir:
618+
output_dir = Path(tmpdir)
619+
validation_dir = output_dir / "validation_images"
620+
validation_dir.mkdir()
621+
(validation_dir / "step_50_prompt_0_64x64.jpg").write_bytes(b"jpg")
622+
(validation_dir / "step_50_prompt_adapter_0_64x64.jpg").write_bytes(b"adapter")
623+
(validation_dir / "step_49_prompt_0_64x64.jpg").write_bytes(b"old")
624+
(validation_dir / "step_50_prompt_0.wav").write_bytes(b"wav")
625+
626+
hub_manager = object.__new__(HubManager)
627+
hub_manager.config = SimpleNamespace(output_dir=str(output_dir))
628+
629+
images, audios = hub_manager._filter_checkpoint_validation_media(
630+
50,
631+
{"prompt": [], "prompt_adapter": []},
632+
{"prompt": []},
633+
)
634+
635+
self.assertEqual(
636+
{key: [Path(path).name for path in paths] for key, paths in images.items()},
637+
{
638+
"prompt": ["step_50_prompt_0_64x64.jpg"],
639+
"prompt_adapter": ["step_50_prompt_adapter_0_64x64.jpg"],
640+
},
641+
)
642+
self.assertEqual(
643+
{key: [Path(path).name for path in paths] for key, paths in audios.items()},
644+
{"prompt": ["step_50_prompt_0.wav"]},
645+
)
646+
564647
def test_save_training_config_sanitizes_public_export(self):
565648
config = SimpleNamespace(
566649
output_dir="output/test",

0 commit comments

Comments
 (0)