Gradio WebUI installieren

Last modified by René Schmidt on 2025/04/02 14:40

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.plotthumb_up
    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)