From 33ab4d13d1fc82e4cf47ffebece10b3f6cc5aa0e Mon Sep 17 00:00:00 2001 From: ndrw1221 Date: Thu, 7 Aug 2025 07:20:37 +0000 Subject: [PATCH 1/2] Update seconds_start based on latent crop in PreEncodedDataset --- stable_audio_tools/data/dataset.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/stable_audio_tools/data/dataset.py b/stable_audio_tools/data/dataset.py index 7543ac17..86cc4d15 100644 --- a/stable_audio_tools/data/dataset.py +++ b/stable_audio_tools/data/dataset.py @@ -319,8 +319,15 @@ def __getitem__(self, idx): start = random.randint(0, last_ix - self.latent_crop_length) else: start = 0 - - latents = latents[:, start:start+self.latent_crop_length] + + # Update seconds_start based on latent crop + original_length = info["seconds_total"] * ( + info["timestamps"][1] - info["timestamps"][0] + ) + seconds_per_latent = original_length / info["padding_mask"].count(1) + info["seconds_start"] += start * seconds_per_latent + + latents = latents[:, start : start + self.latent_crop_length] info["padding_mask"] = info["padding_mask"][start:start+self.latent_crop_length] From 53ec3bf1163e2ca11c76126c435614c6e844e265 Mon Sep 17 00:00:00 2001 From: ndrw1221 Date: Thu, 7 Aug 2025 10:34:30 +0000 Subject: [PATCH 2/2] Update seconds_start calculation in PreEncodedDataset to use floor value --- stable_audio_tools/data/dataset.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/stable_audio_tools/data/dataset.py b/stable_audio_tools/data/dataset.py index 86cc4d15..199f82a5 100644 --- a/stable_audio_tools/data/dataset.py +++ b/stable_audio_tools/data/dataset.py @@ -11,6 +11,7 @@ import torch import torchaudio import webdataset as wds +import math from os import path from torch import nn @@ -325,7 +326,7 @@ def __getitem__(self, idx): info["timestamps"][1] - info["timestamps"][0] ) seconds_per_latent = original_length / info["padding_mask"].count(1) - info["seconds_start"] += start * seconds_per_latent + info["seconds_start"] += math.floor(start * seconds_per_latent) latents = latents[:, start : start + self.latent_crop_length]