Wiki-Quellcode von Gradio WebUI installieren
Zuletzt geändert von René Schmidt am 2025/04/02 14:40
Zeige letzte Bearbeiter
| author | version | line-number | content |
|---|---|---|---|
| 1 | pip install matplotlib | ||
| 2 | |||
| 3 | app.py anlegen | ||
| 4 | |||
| 5 | Code einfügen: | ||
| 6 | |||
| 7 | import torch | ||
| 8 | import gradio as gr | ||
| 9 | import librosa | ||
| 10 | import numpy as np | ||
| 11 | from model import DiT, CFM | ||
| 12 | from huggingface_hub import hf_hub_download | ||
| 13 | from mutagen.mp3 import MP3 | ||
| 14 | import matplotlib.pyplot as plt | ||
| 15 | |||
| 16 | # Device bestimmen (GPU oder CPU) | ||
| 17 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | ||
| 18 | |||
| 19 | # Bereite das Modell vor | ||
| 20 | def prepare_model(device, repo_id="ASLP-lab/DiffRhythm-base"): | ||
| 21 | dit_ckpt_path = hf_hub_download(repo_id=repo_id, filename="cfm_model.pt", cache_dir="./pretrained") | ||
| 22 | dit_config_path = "./config/diffrhythm-1b.json" | ||
| 23 | with open(dit_config_path) as f: | ||
| 24 | model_config = json.load(f) | ||
| 25 | dit_model_cls = DiT | ||
| 26 | cfm = CFM( | ||
| 27 | transformer=dit_model_cls(~*~*model_config["model"]), | ||
| 28 | num_channels=model_config["model"]['mel_dim'] | ||
| 29 | ) | ||
| 30 | cfm = cfm.to(device) | ||
| 31 | cfm = load_checkpoint(cfm, dit_ckpt_path, device=device, use_ema=False) | ||
| 32 | |||
| 33 | tokenizer = CNENTokenizer() | ||
| 34 | muq = MuQMuLan.from_pretrained("OpenMuQ/MuQ-MuLan-large", cache_dir="./pretrained") | ||
| 35 | muq = muq.to(device).eval() | ||
| 36 | |||
| 37 | vae_ckpt_path = hf_hub_download(repo_id="ASLP-lab/DiffRhythm-vae", filename="vae_model.pt", cache_dir="./pretrained") | ||
| 38 | vae = torch.jit.load(vae_ckpt_path, map_location='cpu').to(device) | ||
| 39 | |||
| 40 | return cfm, tokenizer, muq, vae | ||
| 41 | |||
| 42 | # Hauptfunktion für Audio-Stil-Inferenz | ||
| 43 | def infer_audio_style(lrc_text, audio_path, style_text): | ||
| 44 | cfm, tokenizer, muq, vae = prepare_model(device) | ||
| 45 | |||
| 46 | # Extrahiere das Style Prompt von Audio und Text | ||
| 47 | lrc_emb, _ = get_lrc_token(lrc_text, tokenizer, device) | ||
| 48 | style_emb = get_style_prompt(muq, audio_path) | ||
| 49 | |||
| 50 | # Beispielhafte Funktionalität für die Audioerzeugung | ||
| 51 | audio_output = "Audio wurde mit dem Stil erzeugt." | ||
| 52 | |||
| 53 | # Speichere die Audio-Datei (z. B. als WAV-Datei) | ||
| 54 | output_audio_path = "output/generated_audio.wav" | ||
| 55 | # Hier kannst du den Code zur tatsächlichen Audioerzeugung hinzufügen | ||
| 56 | |||
| 57 | # Wellenform visualisieren | ||
| 58 | plot_waveform(output_audio_path) | ||
| 59 | |||
| 60 | return audio_output, output_audio_path | ||
| 61 | |||
| 62 | # Wellenform anzeigen | ||
| 63 | def plot_waveform(audio_path): | ||
| 64 | y, sr = librosa.load(audio_path, sr=None) | ||
| 65 | plt.figure(figsize=(10, 4)) | ||
| 66 | plt.plot(y) | ||
| 67 | plt.title("Audio Wellenform") | ||
| 68 | plt.xlabel("Samples") | ||
| 69 | plt.ylabel("Amplitude") | ||
| 70 | |||
| 71 | waveform_img_path = "output/waveform.png" | ||
| 72 | plt.savefig(waveform_img_path) | ||
| 73 | plt.close() | ||
| 74 | |||
| 75 | # Gradio Interface erstellen | ||
| 76 | interface = gr.Interface( | ||
| 77 | fn=infer_audio_style, | ||
| 78 | inputs=[ | ||
| 79 | gr.Textbox(label="Lyrics mit Zeitmarken (LRC)", placeholder="Gib hier den Text ein..."), | ||
| 80 | gr.Audio(label="Audio-Datei", type="filepath"), | ||
| 81 | gr.Textbox(label="Stiltext", placeholder="Gib den Stiltext ein..."), | ||
| 82 | ], | ||
| 83 | outputs=[ | ||
| 84 | gr.Textbox(label="Generierte Ausgabe"), | ||
| 85 | gr.File(label="Generierte Audio-Datei") # Ohne 'file_path' Parameter | ||
| 86 | ], | ||
| 87 | ) | ||
| 88 | |||
| 89 | # Starte das Interface und mache es im gesamten Netzwerk verfügbar | ||
| 90 | interface.launch(server_name="0.0.0.0", server_port=7860, share=False) | ||
| 91 | |||
| 92 |