Can't train with fp16 on Nvidia P100 #15
Description
Activity
the problem does not occur on torch version 1.6.0 as in the requirements.txt
Use
nvcr.io/nvidia/pytorch:20.07-pydocker image.I have the same issue, any idea of where the complex number is generated?,
(for 1.6 it works fine but I want to combine this with another code that requires pytorch 1.9)I got the same issue. It’s due to a bug in the pytorch STFT function for half tensor. The work around is moving the calculation of y_hat_mel in train.py outside autocast, and casting y_hat to float one line above y_hat_mel calculation.
@boltzmann-Li Can you create a PR so we can see that fix. I haven't managed to get it working following your instructions.
FYI the problem hasn't been fixed in torch 1.10.0
Is there an issue for the Complex Half problem?
@boltzmann-Li Can you create a PR so we can see that fix. I haven't managed to get it working following your instructions.
FYI the problem hasn't been fixed in torch 1.10.0
Is there an issue for the Complex Half problem?
I created a pull request. It has been working for me with 3090 GPUs and torch 1.9
Reacted by Harry Coultas Blum, Faris Hijazi, Đỗ Trí Nhân, AndreyBocharnikov, splinter21, akfheaven, barryhunt and WALKERReacted by Faris Hijazi and akfheavenReacted by Faris Hijazi and akfheavenReacted by Faris Hijazi, AndreyBocharnikov and akfheavenReacted by Faris Hijazi and akfheavenVery helpful @boltzmann-Li
here are the lines https://github.com/boltzmann-Li/vits/blob/5a1f4b7afb8a822f66c0ddc75bc959a44a57d035/train_ms.py#L156-L166
Reacted by Đỗ Trí Nhân, Hansss, hdmjdp, barryhunt, Shuangsheng Duo and Yoga Tiara WigunaI think a better way to solve this problem is to wrap the torch.stft with
autocast(enabled=off)inside the mel_spectrogram_torch function. Here is the code:def mel_spectrogram_torch(y, n_fft, num_mels, sampling_rate, hop_size, win_size, fmin, fmax, center=False): if torch.min(y) < -1.: print('min value is ', torch.min(y)) if torch.max(y) > 1.: print('max value is ', torch.max(y)) global mel_basis, hann_window dtype_device = str(y.dtype) + '_' + str(y.device) fmax_dtype_device = str(fmax) + '_' + dtype_device wnsize_dtype_device = str(win_size) + '_' + dtype_device if fmax_dtype_device not in mel_basis: mel = librosa_mel_fn(sampling_rate, n_fft, num_mels, fmin, fmax) mel_basis[fmax_dtype_device] = torch.from_numpy(mel).to(dtype=y.dtype, device=y.device) if wnsize_dtype_device not in hann_window: hann_window[wnsize_dtype_device] = torch.hann_window(win_size).to(dtype=y.dtype, device=y.device) y = torch.nn.functional.pad(y.unsqueeze(1), (int((n_fft-hop_size)/2), int((n_fft-hop_size)/2)), mode='reflect') y = y.squeeze(1) with autocast(enabled=False): y = y.float() spec = torch.stft(y, n_fft, hop_length=hop_size, win_length=win_size, window=hann_window[wnsize_dtype_device], center=center, pad_mode='reflect', normalized=False, onesided=True) spec = torch.sqrt(spec.pow(2).sum(-1) + 1e-6) spec = torch.matmul(mel_basis[fmax_dtype_device], spec) spec = spectral_normalize_torch(spec) return specReacted by Harry Coultas Blum, Daniel Doña, Hansss, hdmjdp, GuanXuzeng, Will Rice, Tmn07, wuhc, Alan Sun, Yiwei Guo and 9 moreReacted by GuanXuzeng and Alan SunReacted by Alan Sun and ConsistencyVC
training with fp16 doesn't work for me on a P100, I'll look into fixing it, but for future reference here is the full stacktrace
torch version 1.9.0