Wire total_weights across trainers and simplify metric mappings
- Wire total_weights into telemetry: Pass the existing `num_model_parameters` calculation into `record_scalar_metrics(..., total_weights=num_model_parameters)` across all trainers (Stable Diffusion, SDXL, Flux, Wan, and DreamBooth). This populates the total_weights card in Google Cloud ML Diagnostics. - Count text-encoder parameters when they are trainable: Stable Diffusion and DreamBooth apply gradients to text_encoder_state when train_text_encoder is enabled, so its parameters are now added to num_model_parameters. The calculation stays outside the training loop. SDXL, Flux, and Wan reject train_text_encoder in __init__, so they are unaffected. - Remove the redundant `metric_types` import and `if/else` branching. Standardize `_METRICS_TO_MANAGED` directly on canonical string literals, matching the SDK's internal representation. - Add stable_diffusion_trainer_test and dreambooth_trainer_test covering both train_text_encoder values, and document custom metrics in docs/metrics.md.
R
Richa Gupta committed
043f62be8f0a042221d51809ca2d6fc2a593900c
Parent: 1bc5481