-
Notifications
You must be signed in to change notification settings - Fork 30
Expand file tree
/
Copy pathsde_denoise_curve.py
More file actions
114 lines (92 loc) · 3.76 KB
/
Copy pathsde_denoise_curve.py
File metadata and controls
114 lines (92 loc) · 3.76 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
#!/usr/bin/env python3
"""SDE denoise curve workflow: per-frame re-noise modulation.
Replaces test_denoise_curve_E.py. Demonstrates:
- DiffusionConfig with method="sde"
- CurveRamp feeding sde_denoise_curve on Generate
- High curve values get full re-noise (more transformation),
low values pull toward source (preserve source)
"""
import os
import sys
import time
import soundfile as sf
import torch
project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))
if project_root not in sys.path:
sys.path.insert(0, project_root)
from acestep.nodes import Audio
from acestep.nodes.model_nodes import LoadModel
from acestep.nodes.vae_nodes import VAEEncodeAudio, VAEDecodeAudio
from acestep.nodes.cond_nodes import TextEncode
from acestep.nodes.semantic_nodes import SemanticExtract
from acestep.nodes.curve_nodes import CurveRamp
from acestep.nodes.diffusion_nodes import DiffusionConfigNode, Generate
from acestep.constants import TASK_INSTRUCTIONS
from acestep.fixtures import audio_fixture
SOURCE_AUDIO = str(audio_fixture("inside_confusion_loop_60s_gsm.wav"))
OUTPUT_DIR = os.path.join(project_root, "test_output", "examples")
def load_audio(path: str, duration: float = 60.0) -> Audio:
data, sr = sf.read(path, dtype="float32")
waveform = torch.from_numpy(data.T if data.ndim > 1 else data.reshape(1, -1))
if sr != 48000:
import torchaudio
waveform = torchaudio.transforms.Resample(sr, 48000)(waveform)
waveform = waveform[:2, :int(duration * 48000)]
return Audio(waveform=waveform, sample_rate=48000)
def save_audio(audio: Audio, path: str) -> None:
wav = audio.waveform
if wav.dim() == 3:
wav = wav.squeeze(0)
sf.write(path, wav.cpu().numpy().T, audio.sample_rate)
print(f"Saved: {path}")
def main():
print("=" * 70)
print("WORKFLOW: SDE Denoise Curve (per-frame re-noise modulation)")
print(" Ramp 0.3 -> 1.0 across 60s")
print("=" * 70)
os.makedirs(OUTPUT_DIR, exist_ok=True)
# --- Load model ---
handles = LoadModel().execute(
project_root=project_root,
config_path="acestep-v15-turbo",
device="cuda",
use_flash_attention=True,
)
model, clip, vae = handles["model"], handles["clip"], handles["vae"]
# --- Encode source ---
source_audio = load_audio(SOURCE_AUDIO)
source_latent = VAEEncodeAudio().execute(vae=vae, audio=source_audio)["latent"]
T = source_latent.tensor.shape[1]
context_latent = SemanticExtract().execute(model=model, latent=source_latent)["latent"]
conditioning = TextEncode().execute(
clip=clip, model=model,
refer_latent=source_latent,
tags="deathstep death deaht deaht",
instruction=TASK_INSTRUCTIONS["cover"],
bpm=136, duration=60.0, key="G# minor",
)["conditioning"]
# --- SDE denoise curve: 0.3 (preserve start) -> 1.0 (transform end) ---
denoise_curve = CurveRamp().execute(
start=0.3, end=1.0, length=T,
)["curve"]
print(f"SDE denoise curve: {denoise_curve.tensor[0]:.2f} -> {denoise_curve.tensor[-1]:.2f}")
# --- Generate with SDE method + denoise curve ---
config = DiffusionConfigNode().execute(
steps=8, shift=3.0, seed=1528, method="sde",
)["config"]
t0 = time.time()
output_latent = Generate().execute(
model=model,
config=config,
positive=conditioning,
context_latent=context_latent,
source_latent=source_latent,
sde_denoise_curve=denoise_curve,
)["latent"]
print(f"Generated in {time.time() - t0:.2f}s")
# --- Decode ---
output_audio = VAEDecodeAudio().execute(vae=vae, latent=output_latent)["audio"]
save_audio(output_audio, os.path.join(OUTPUT_DIR, "sde_denoise_curve.wav"))
print("\nDone.")
if __name__ == "__main__":
main()