SIGN IN SIGN UP

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