@@ -374,7 +374,8 @@ class MelScale(torch.nn.Module):
374374 :py:func:`torchaudio.functional.melscale_fbanks` - The function used to
375375 generate the filter banks.
376376 """
377- __constants__ = ["n_mels" , "sample_rate" , "f_min" , "f_max" ]
377+
378+ __constants__ = ["n_mels" , "sample_rate" , "f_min" , "f_max" , "n_stft" ]
378379
379380 def __init__ (
380381 self ,
@@ -391,13 +392,16 @@ def __init__(
391392 self .sample_rate = sample_rate
392393 self .f_max = f_max if f_max is not None else float (sample_rate // 2 )
393394 self .f_min = f_min
395+ self .n_stft = n_stft
394396 self .norm = norm
395397 self .mel_scale = mel_scale
396398
397399 if f_min > self .f_max :
398400 raise ValueError ("Require f_min: {} <= f_max: {}" .format (f_min , self .f_max ))
399401
400- fb = F .melscale_fbanks (n_stft , self .f_min , self .f_max , self .n_mels , self .sample_rate , self .norm , self .mel_scale )
402+ fb = F .melscale_fbanks (
403+ self .n_stft , self .f_min , self .f_max , self .n_mels , self .sample_rate , self .norm , self .mel_scale
404+ )
401405 self .register_buffer ("fb" , fb )
402406
403407 def forward (self , specgram : Tensor ) -> Tensor :
@@ -464,10 +468,13 @@ def __init__(
464468 driver : str = "gels" ,
465469 ) -> None :
466470 super (InverseMelScale , self ).__init__ ()
471+ self .n_stft = n_stft
467472 self .n_mels = n_mels
468473 self .sample_rate = sample_rate
469474 self .f_max = f_max or float (sample_rate // 2 )
470475 self .f_min = f_min
476+ self .norm = norm
477+ self .mel_scale = mel_scale
471478 self .driver = driver
472479
473480 if f_min > self .f_max :
@@ -476,7 +483,9 @@ def __init__(
476483 if driver not in ["gels" , "gelsy" , "gelsd" , "gelss" ]:
477484 raise ValueError (f'driver must be one of ["gels", "gelsy", "gelsd", "gelss"]. Found { driver } .' )
478485
479- fb = F .melscale_fbanks (n_stft , self .f_min , self .f_max , self .n_mels , self .sample_rate , norm , mel_scale )
486+ fb = F .melscale_fbanks (
487+ self .n_stft , self .f_min , self .f_max , self .n_mels , self .sample_rate , self .norm , self .mel_scale
488+ )
480489 self .register_buffer ("fb" , fb )
481490
482491 def forward (self , melspec : Tensor ) -> Tensor :
0 commit comments