kohya-ss/sd-scripts

Question about V-Prediction in SDXL Finetuning

Aperta

#1163 aperta il 9 mar 2024

 (1 commento) (0 reazioni) (0 assegnatari)Python (1218 fork)batch import
help wanted

Metriche repository

Star
 (7198 stelle)
Metriche merge PR
 (Merge medio 18h 39m) (16 PR mergiate in 30 g)

Descrizione

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

Guida contributor