Gradio WebUI installieren
pip install matplotlib
app.py anlegen
Code einfügen:
import torch
import gradio as gr
import librosa
import numpy as np
from model import DiT, CFM
from huggingface_hub import hf_hub_download
from mutagen.mp3 import MP3
import matplotlib.pyplot as plt
# Device bestimmen (GPU oder CPU)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Bereite das Modell vor
def prepare_model(device, repo_id="ASLP-lab/DiffRhythm-base"):
dit_ckpt_path = hf_hub_download(repo_id=repo_id, filename="cfm_model.pt", cache_dir="./pretrained")
dit_config_path = "./config/diffrhythm-1b.json"
with open(dit_config_path) as f:
model_config = json.load(f)
dit_model_cls = DiT
cfm = CFM(
transformer=dit_model_cls(**model_config["model"]),
num_channels=model_config["model"]['mel_dim']
)
cfm = cfm.to(device)
cfm = load_checkpoint(cfm, dit_ckpt_path, device=device, use_ema=False)
tokenizer = CNENTokenizer()
muq = MuQMuLan.from_pretrained("OpenMuQ/MuQ-MuLan-large", cache_dir="./pretrained")
muq = muq.to(device).eval()
vae_ckpt_path = hf_hub_download(repo_id="ASLP-lab/DiffRhythm-vae", filename="vae_model.pt", cache_dir="./pretrained")
vae = torch.jit.load(vae_ckpt_path, map_location='cpu').to(device)
return cfm, tokenizer, muq, vae
# Hauptfunktion für Audio-Stil-Inferenz
def infer_audio_style(lrc_text, audio_path, style_text):
cfm, tokenizer, muq, vae = prepare_model(device)
# Extrahiere das Style Prompt von Audio und Text
lrc_emb, _ = get_lrc_token(lrc_text, tokenizer, device)
style_emb = get_style_prompt(muq, audio_path)
# Beispielhafte Funktionalität für die Audioerzeugung
audio_output = "Audio wurde mit dem Stil erzeugt."
# Speichere die Audio-Datei (z. B. als WAV-Datei)
output_audio_path = "output/generated_audio.wav"
# Hier kannst du den Code zur tatsächlichen Audioerzeugung hinzufügen
# Wellenform visualisieren
plot_waveform(output_audio_path)
return audio_output, output_audio_path
# Wellenform anzeigen
def plot_waveform(audio_path):
y, sr = librosa.load(audio_path, sr=None)
plt.figure(figsize=(10, 4))
plt.plot![]()
plt.title("Audio Wellenform")
plt.xlabel("Samples")
plt.ylabel("Amplitude")
waveform_img_path = "output/waveform.png"
plt.savefig(waveform_img_path)
plt.close()
# Gradio Interface erstellen
interface = gr.Interface(
fn=infer_audio_style,
inputs=[
gr.Textbox(label="Lyrics mit Zeitmarken (LRC)", placeholder="Gib hier den Text ein..."),
gr.Audio(label="Audio-Datei", type="filepath"),
gr.Textbox(label="Stiltext", placeholder="Gib den Stiltext ein..."),
],
outputs=[
gr.Textbox(label="Generierte Ausgabe"),
gr.File(label="Generierte Audio-Datei") # Ohne 'file_path' Parameter
],
)
# Starte das Interface und mache es im gesamten Netzwerk verfügbar
interface.launch(server_name="0.0.0.0", server_port=7860, share=False)