Skip to content

Commit 26e7661

Browse files
authored
Merge pull request #3140 from bghira/feature/better-music-audio-model-card-data
add XM and NextLat info to model card, plus audio-specific fixes
2 parents 4926fbd + 51e4e8e commit 26e7661

4 files changed

Lines changed: 203 additions & 18 deletions

File tree

simpletuner/helpers/models/common.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6913,6 +6913,12 @@ def custom_model_card_schedule_info(self):
69136913
"""
69146914
return []
69156915

6916+
def custom_model_card_training_mode_info(self, args) -> str:
6917+
"""
6918+
Override this in a subclass to add model-specific training mode details to model cards.
6919+
"""
6920+
return ""
6921+
69166922
def custom_model_card_code_example(self, repo_id: str = None) -> str:
69176923
"""
69186924
Override this to provide custom code examples for model cards.

simpletuner/helpers/models/minimaxmusic/model.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -318,6 +318,18 @@ def __init__(self, config, accelerator):
318318
self.TEXT_ENCODER_CONFIGURATION = {}
319319
self.DEFAULT_LORA_TARGET = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]
320320

321+
def custom_model_card_training_mode_info(self, args) -> str:
322+
train_component = str(getattr(args, "minimax_music_train_component", "transformer") or "transformer")
323+
component_labels = {
324+
"language_model": "language_model (global LM / RVQ planner)",
325+
"transformer": "transformer (DiT/audio denoiser)",
326+
}
327+
lines = [f"- MiniMax Music train component: `{component_labels.get(train_component, train_component)}`"]
328+
lm_max_frames = getattr(args, "minimax_music_lm_max_frames", None)
329+
if lm_max_frames:
330+
lines.append(f"- MiniMax Music LM max frames: `{lm_max_frames}`")
331+
return "\n".join(lines)
332+
321333
@classmethod
322334
def max_swappable_blocks(cls, config=None) -> Optional[int]:
323335
return 35

simpletuner/helpers/publishing/metadata.py

Lines changed: 85 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -433,6 +433,57 @@ def model_card_note(args):
433433
return f"\n**Note:** {note_contents}\n"
434434

435435

436+
def _model_card_bool(args, name: str) -> bool:
437+
value = getattr(args, name, False)
438+
if isinstance(value, bool):
439+
return value
440+
if isinstance(value, (int, float)):
441+
return bool(value)
442+
if isinstance(value, str):
443+
return value.strip().lower() in {"1", "true", "yes", "on"}
444+
return False
445+
446+
447+
def _model_card_optional_value(args, name: str):
448+
value = getattr(args, name, None)
449+
if isinstance(value, (str, int, float, bool)) or value is None:
450+
return value
451+
return None
452+
453+
454+
def _model_card_training_modes(args, model: Optional[ModelFoundation] = None) -> str:
455+
lines = []
456+
model_specific_info = ""
457+
if model and hasattr(model, "custom_model_card_training_mode_info"):
458+
model_specific_info = model.custom_model_card_training_mode_info(args) or ""
459+
if not isinstance(model_specific_info, str):
460+
model_specific_info = ""
461+
model_specific_info = model_specific_info.strip()
462+
463+
if _model_card_bool(args, "nextlat_enabled"):
464+
lines.append("- NextLat: Enabled")
465+
lines.append(f" - Block index: `{_model_card_optional_value(args, 'nextlat_block_index')}`")
466+
lines.append(f" - Weight: `{_model_card_optional_value(args, 'nextlat_weight')}`")
467+
lines.append(f" - State loss: `{_model_card_optional_value(args, 'nextlat_state_loss')}`")
468+
lines.append(f" - KL weight: `{_model_card_optional_value(args, 'nextlat_kl_weight')}`")
469+
470+
if _model_card_bool(args, "xm_enabled"):
471+
lines.append("- XM: Enabled")
472+
lines.append(f" - Candidate count: `{_model_card_optional_value(args, 'xm_candidate_count')}`")
473+
lines.append(f" - Selection scope: `{_model_card_optional_value(args, 'xm_selection_scope')}`")
474+
lines.append(f" - Training target: `{_model_card_optional_value(args, 'xm_training_target')}`")
475+
lines.append(f" - Block size: `{_model_card_optional_value(args, 'xm_block_size')}`")
476+
477+
if not lines and not model_specific_info:
478+
return ""
479+
blocks = []
480+
if model_specific_info:
481+
blocks.append(model_specific_info)
482+
if lines:
483+
blocks.append("\n".join(lines))
484+
return "## Training modes\n\n" + "\n".join(blocks) + "\n\n"
485+
486+
436487
def save_metadata_sample(
437488
image_path: str,
438489
image: Union[Image.Image, np.ndarray, list, str, torch.Tensor],
@@ -772,6 +823,36 @@ def _add_widget_entries(media, asset_prefix: str):
772823
gallery_intro = "You can find some example audio samples in the following gallery:"
773824
else:
774825
gallery_intro = "You can find some example images in the following gallery:"
826+
gallery_section = f"{gallery_intro}\n\n<Gallery />\n" if has_media else ""
827+
validation_disabled = _model_card_bool(args, "validation_disable")
828+
validation_intro = (
829+
"Validation was disabled during training."
830+
if validation_disabled
831+
else (
832+
"The main validation prompt used during training was:"
833+
if prompt
834+
else (
835+
"Validation used ground-truth images as an input for partial denoising (img2img)."
836+
if args.validation_using_datasets
837+
else "No validation prompt was used during training."
838+
)
839+
)
840+
)
841+
validation_prompt_block = f"```\n{prompt}\n```\n" if prompt and not validation_disabled else ""
842+
validation_settings = ""
843+
if not validation_disabled:
844+
validation_settings = f"""## Validation settings
845+
- CFG: `{StateTracker.get_args().validation_guidance}`
846+
- CFG Rescale: `{StateTracker.get_args().validation_guidance_rescale}`
847+
- Steps: `{StateTracker.get_args().validation_num_inference_steps}`
848+
- Sampler: `{_validation_scheduler_label(model, StateTracker.get_args())}`
849+
- Seed: `{StateTracker.get_args().validation_seed}`
850+
- Resolution{'s' if ',' in str(StateTracker.get_args().validation_resolution) else ''}: `{str(StateTracker.get_args().validation_resolution)}`
851+
{f"- Skip-layer guidance: {_skip_layers(args)}" if args.model_family in ['sd3', 'flux'] else ''}
852+
853+
Note: The validation settings are not necessarily the same as the [training settings](#training-settings).
854+
855+
"""
775856
sage_usage = getattr(args.sageattention_usage, "value", args.sageattention_usage)
776857
license_metadata = _license_metadata(model)
777858
yaml_content = f"""---
@@ -800,26 +881,11 @@ def _add_widget_entries(media, asset_prefix: str):
800881
801882
This is a {model_type(args)} derived from [{base_model}](https://huggingface.co/{base_model}).
802883
803-
{'The main validation prompt used during training was:' if prompt else 'Validation used ground-truth images as an input for partial denoising (img2img).' if args.validation_using_datasets else 'No validation prompt was used during training.'}
804-
{'```' if prompt else ''}
805-
{prompt}
806-
{'```' if prompt else ''}
884+
{validation_intro}
885+
{validation_prompt_block}
807886
808887
{model_card_note(args)}
809-
## Validation settings
810-
- CFG: `{StateTracker.get_args().validation_guidance}`
811-
- CFG Rescale: `{StateTracker.get_args().validation_guidance_rescale}`
812-
- Steps: `{StateTracker.get_args().validation_num_inference_steps}`
813-
- Sampler: `{_validation_scheduler_label(model, StateTracker.get_args())}`
814-
- Seed: `{StateTracker.get_args().validation_seed}`
815-
- Resolution{'s' if ',' in str(StateTracker.get_args().validation_resolution) else ''}: `{str(StateTracker.get_args().validation_resolution)}`
816-
{f"- Skip-layer guidance: {_skip_layers(args)}" if args.model_family in ['sd3', 'flux'] else ''}
817-
818-
Note: The validation settings are not necessarily the same as the [training settings](#training-settings).
819-
820-
{gallery_intro}\n
821-
822-
<Gallery />
888+
{validation_settings}{gallery_section}
823889
824890
The text encoder {'**was**' if train_text_encoder else '**was not**'} trained.
825891
{'You may reuse the base model text encoder for inference.' if not train_text_encoder else 'If the text encoder from this repository is not used at inference time, unexpected or bad results could occur.'}
@@ -849,6 +915,7 @@ def _add_widget_entries(media, asset_prefix: str):
849915
{('- SLA: Enabled (you MUST use SLA for inference; sla_attention.pt contains attention weights)') if StateTracker.get_args().attention_mechanism == 'sla' else ''}
850916
{lora_info(args=StateTracker.get_args())}
851917
918+
{_model_card_training_modes(args, model=model)}
852919
## Datasets
853920
854921
{datasets_str}

tests/test_model_card.py

Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,7 @@ def setUp(self):
4141
self.args.model_type = "lora"
4242
self.args.model_family = "sdxl"
4343
self.args.validation_prompt = "A test prompt"
44+
self.args.validation_disable = False
4445
self.args.validation_negative_prompt = "A negative prompt"
4546
self.args.validation_num_inference_steps = 50
4647
self.args.validation_guidance = 7.5
@@ -84,6 +85,18 @@ def setUp(self):
8485
self.args.t5_padding = "unmodified"
8586
self.args.enable_xformers_memory_efficient_attention = False
8687
self.args.attention_mechanism = "diffusers"
88+
self.args.minimax_music_train_component = None
89+
self.args.minimax_music_lm_max_frames = None
90+
self.args.nextlat_enabled = False
91+
self.args.nextlat_block_index = -1
92+
self.args.nextlat_weight = 0.0
93+
self.args.nextlat_state_loss = "smooth_l1"
94+
self.args.nextlat_kl_weight = 0.0
95+
self.args.xm_enabled = False
96+
self.args.xm_candidate_count = 1
97+
self.args.xm_selection_scope = "sample"
98+
self.args.xm_training_target = "noise"
99+
self.args.xm_block_size = 0
87100
self.mock_model = MagicMock(MODEL_TYPE=MagicMock(value="unet"))
88101

89102
def test_model_imports(self):
@@ -293,6 +306,93 @@ def test_audio_model_families_use_audio_pipeline_tag(self):
293306
self.assertEqual(_pipeline_tag(self.args), "text-to-audio")
294307
self.assertEqual(_secondary_pipeline_tag(self.args), "audio")
295308

309+
def test_minimax_music_model_card_reports_modes_and_disabled_validation(self):
310+
self.args.model_family = "minimaxmusic"
311+
self.args.model_type = "lora"
312+
self.args.lora_type = "standard"
313+
self.args.peft_lora_mode = "standard"
314+
self.args.controlnet = False
315+
self.args.control = False
316+
self.args.validation_disable = True
317+
self.args.model_card_note = ""
318+
self.args.minimax_music_train_component = "language_model"
319+
self.args.minimax_music_lm_max_frames = 128
320+
self.args.nextlat_enabled = True
321+
self.args.nextlat_block_index = -1
322+
self.args.nextlat_weight = 0.1
323+
self.args.nextlat_state_loss = "smooth_l1"
324+
self.args.nextlat_kl_weight = 0.0
325+
self.args.xm_enabled = True
326+
self.args.xm_candidate_count = 2
327+
self.args.xm_selection_scope = "block"
328+
self.args.xm_training_target = "route"
329+
self.args.xm_block_size = 16
330+
model = MagicMock(
331+
MODEL_LICENSE="other",
332+
PREDICTION_TYPE=SimpleNamespace(value="autoregressive_next_token"),
333+
gligen=False,
334+
)
335+
model.validation_audio_sample_rate.return_value = 44100
336+
model.custom_model_card_schedule_info.return_value = ""
337+
model.custom_model_card_code_example.return_value = "```python\npass\n```"
338+
model.custom_model_card_training_mode_info.return_value = (
339+
"- MiniMax Music train component: `language_model (global LM / RVQ planner)`\n"
340+
"- MiniMax Music LM max frames: `128`"
341+
)
342+
343+
with tempfile.TemporaryDirectory() as tmpdir:
344+
with (
345+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_model_family", return_value="minimaxmusic"),
346+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_data_backends", return_value={}),
347+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_weight_dtype", return_value=torch.bfloat16),
348+
patch(
349+
"simpletuner.helpers.publishing.metadata.StateTracker.get_accelerator",
350+
return_value=MagicMock(num_processes=1),
351+
),
352+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_args", return_value=self.args),
353+
patch("simpletuner.helpers.publishing.metadata.StateTracker.get_model", return_value=model),
354+
):
355+
save_model_card(
356+
repo_id="test-repo",
357+
images=None,
358+
audios=None,
359+
base_model="MiniMaxAI/MiniMax-Music3",
360+
train_text_encoder=False,
361+
prompt="",
362+
validation_prompts=None,
363+
validation_shortnames=None,
364+
repo_folder=tmpdir,
365+
model=model,
366+
global_step=6000,
367+
epoch=250,
368+
)
369+
370+
readme = Path(tmpdir, "README.md").read_text(encoding="utf-8")
371+
self.assertIn("Validation was disabled during training.", readme)
372+
self.assertNotIn("## Validation settings", readme)
373+
self.assertNotIn("<Gallery />", readme)
374+
self.assertIn("## Training modes", readme)
375+
self.assertIn("- MiniMax Music train component: `language_model (global LM / RVQ planner)`", readme)
376+
self.assertIn("- MiniMax Music LM max frames: `128`", readme)
377+
self.assertIn("- NextLat: Enabled", readme)
378+
self.assertIn(" - Weight: `0.1`", readme)
379+
self.assertIn("- XM: Enabled", readme)
380+
self.assertIn(" - Candidate count: `2`", readme)
381+
self.assertIn(" - Training target: `route`", readme)
382+
model.custom_model_card_training_mode_info.assert_called_once_with(self.args)
383+
384+
def test_minimax_music_model_card_training_mode_info(self):
385+
from simpletuner.helpers.models.minimaxmusic.model import MiniMaxMusic
386+
387+
self.args.minimax_music_train_component = "language_model"
388+
self.args.minimax_music_lm_max_frames = 128
389+
model = MiniMaxMusic.__new__(MiniMaxMusic)
390+
391+
details = model.custom_model_card_training_mode_info(self.args)
392+
393+
self.assertIn("- MiniMax Music train component: `language_model (global LM / RVQ planner)`", details)
394+
self.assertIn("- MiniMax Music LM max frames: `128`", details)
395+
296396
def test_hub_commit_message_omits_diffusion_schedule_fields_for_flow_matching(self):
297397
hub_manager = object.__new__(HubManager)
298398
hub_manager.collected_data_backend_str = "['alt-embed-cache', 'h3-drift0-anyflow-openvid-39f-480']"

0 commit comments

Comments
 (0)