Skip to content

Commit 96200d4

Browse files
authored
Retain init args as attributes in MelScale and InverseMelScale (pytorch#4126)
1 parent e284e58 commit 96200d4

1 file changed

Lines changed: 12 additions & 3 deletions

File tree

src/torchaudio/transforms/_transforms.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)