kohya-ss/sd-scripts

Question about V-Prediction in SDXL Finetuning

Open

#1,163 opened on Mar 9, 2024

 (1 comment) (0 reactions) (0 assignees)Python (1,218 forks)batch import
help wanted

Repository metrics

Stars
 (7,198 stars)
PR merge metrics
 (Avg merge 18h 39m) (16 merged PRs in 30d)

Description

It's just my one-sided doubts, about the implement of the v-prediction. In sdxl training, the source code implements v-prediction by:

def add_v_prediction_like_loss(loss, timesteps, noise_scheduler, v_pred_like_loss):
    scale = get_snr_scale(timesteps, noise_scheduler)
    # print(f"add v-prediction like loss: {v_pred_like_loss}, scale: {scale}, loss: {loss}, time: {timesteps}")
    loss = loss + loss / scale * v_pred_like_loss
    return loss

which is mathematically equivalent to: L:=L+snr*L*w, where w=v_pred_like_loss, and snr=scale, while the paper suggests: L:=snr*L.

So, is the source adds additional v-pred like loss rather than scaling it? Why are the implementation and paper different? I'm not a mathematician, and maybe I'm short-sighted. Hope someone can answer my doubts :D

Contributor guide