CoolFace
Apppublic

parson/audioEditing

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
app.py336 linesDownload Raw Back to root
1import gradio as gr2import random3import torch4import os5from torch import inference_mode6from tempfile import NamedTemporaryFile7import numpy as np8from models import load_model9import utils10from inversion_utils import inversion_forward_process, inversion_reverse_process11 12 13# current_loaded_model = "cvssp/audioldm2-music"14# # current_loaded_model = "cvssp/audioldm2-music"15 16# ldm_stable = load_model(current_loaded_model, device, 200)  # deafult model17LDM2 = "cvssp/audioldm2"18MUSIC = "cvssp/audioldm2-music"19LDM2_LARGE = "cvssp/audioldm2-large"20device = torch.device("cuda" if torch.cuda.is_available() else "cpu")21ldm2 = load_model(model_id=LDM2, device=device)22ldm2_large = load_model(model_id=LDM2_LARGE, device=device)23ldm2_music = load_model(model_id=MUSIC, device=device)24 25 26def randomize_seed_fn(seed, randomize_seed):27    if randomize_seed:28        seed = random.randint(0, np.iinfo(np.int32).max)29    torch.manual_seed(seed)30    return seed31 32 33def invert(ldm_stable, x0, prompt_src, num_diffusion_steps, cfg_scale_src):  # , ldm_stable):34    ldm_stable.model.scheduler.set_timesteps(num_diffusion_steps, device=device)35 36    with inference_mode():37        w0 = ldm_stable.vae_encode(x0)38 39    # find Zs and wts - forward process40    _, zs, wts = inversion_forward_process(ldm_stable, w0, etas=1,41                                           prompts=[prompt_src],42                                           cfg_scales=[cfg_scale_src],43                                           prog_bar=True,44                                           num_inference_steps=num_diffusion_steps,45                                           numerical_fix=True)46    return zs, wts47 48 49def sample(ldm_stable, zs, wts, steps, prompt_tar, tstart, cfg_scale_tar):  # , ldm_stable):50    # reverse process (via Zs and wT)51    tstart = torch.tensor(tstart, dtype=torch.int)52    skip = steps - tstart53    w0, _ = inversion_reverse_process(ldm_stable, xT=wts, skips=steps - skip,54                                      etas=1., prompts=[prompt_tar],55                                      neg_prompts=[""], cfg_scales=[cfg_scale_tar],56                                      prog_bar=True,57                                      zs=zs[:int(steps - skip)])58 59    # vae decode image60    with inference_mode():61        x0_dec = ldm_stable.vae_decode(w0)62    if x0_dec.dim() < 4:63        x0_dec = x0_dec[None, :, :, :]64 65    with torch.no_grad():66        audio = ldm_stable.decode_to_mel(x0_dec)67 68    return (16000, audio.squeeze().cpu().numpy())69 70 71def edit(cache_dir,72         input_audio,73         model_id: str,74         do_inversion: bool,75         wtszs_file: str,76         #  wts: gr.State, zs: gr.State,77         saved_inv_model: str,78         source_prompt="",79         target_prompt="",80         steps=200,81         cfg_scale_src=3.5,82         cfg_scale_tar=12,83         t_start=45,84         randomize_seed=True):85 86    print(model_id)87    if model_id == LDM2:88        ldm_stable = ldm289    elif model_id == LDM2_LARGE:90        ldm_stable = ldm2_large91    else:  # MUSIC92        ldm_stable = ldm2_music93 94    # If the inversion was done for a different model, we need to re-run the inversion95    if not do_inversion and (saved_inv_model is None or saved_inv_model != model_id):96        do_inversion = True97 98    if input_audio is None:99        raise gr.Error('Input audio missing!')100    x0 = utils.load_audio(input_audio, ldm_stable.get_fn_STFT(), device=device)101 102    if not (do_inversion or randomize_seed):103        if not os.path.exists(wtszs_file):104            do_inversion = True105            # Too much time has passed106 107    if do_inversion or randomize_seed:  # always re-run inversion108        zs_tensor, wts_tensor = invert(ldm_stable=ldm_stable, x0=x0, prompt_src=source_prompt,109                                       num_diffusion_steps=steps,110                                       cfg_scale_src=cfg_scale_src)111        f = NamedTemporaryFile("wb", dir=cache_dir, suffix=".pth", delete=False)112        torch.save({'wts': wts_tensor, 'zs': zs_tensor}, f.name)113        wtszs_file = f.name114        # wtszs_file = gr.State(value=f.name)115        # wts = gr.State(value=wts_tensor)116        # zs = gr.State(value=zs_tensor)117        # demo.move_resource_to_block_cache(f.name)118        saved_inv_model = model_id119        do_inversion = False120    else:121        wtszs = torch.load(wtszs_file, map_location=device)122        # wtszs = torch.load(wtszs_file.f, map_location=device)123        wts_tensor = wtszs['wts']124        zs_tensor = wtszs['zs']125 126    # make sure t_start is in the right limit127    # t_start = change_tstart_range(t_start, steps)128 129    output = sample(ldm_stable, zs_tensor, wts_tensor, steps, prompt_tar=target_prompt,130                    tstart=int(t_start / 100 * steps), cfg_scale_tar=cfg_scale_tar)131 132    return output, wtszs_file, saved_inv_model, do_inversion133 134 135def get_example():136    case = [137        ['Examples/Beethoven.wav',138         '',139         'A recording of an arcade game soundtrack.',140         45,141         'cvssp/audioldm2-music',142         '27s',143         'Examples/Beethoven_arcade.wav',144         ],145        ['Examples/Beethoven.wav',146         'A high quality recording of wind instruments and strings playing.',147         'A high quality recording of a piano playing.',148         45,149         'cvssp/audioldm2-music',150         '27s',151         'Examples/Beethoven_piano.wav',152         ],153        ['Examples/ModalJazz.wav',154         'Trumpets playing alongside a piano, bass and drums in an upbeat old-timey cool jazz song.',155         'A banjo playing alongside a piano, bass and drums in an upbeat old-timey cool country song.',156         45,157         'cvssp/audioldm2-music',158         '106s',159         'Examples/ModalJazz_banjo.wav',],160        ['Examples/Cat.wav',161         '',162         'A dog barking.',163         75,164         'cvssp/audioldm2-large',165         '10s',166         'Examples/Cat_dog.wav',]167    ]168    return case169 170 171intro = """172<h1 style="font-weight: 1400; text-align: center; margin-bottom: 7px;"> ZETA Editing 🎧 </h1>173<h2 style="font-weight: 1400; text-align: center; margin-bottom: 7px;"> Zero-Shot Text-Based Audio Editing Using DDPM Inversion πŸŽ›οΈ </h2>174<h3 style="margin-bottom: 10px; text-align: center;">175    <a href="https://arxiv.org/abs/2402.10009">[Paper]</a>&nbsp;|&nbsp;176    <a href="https://hilamanor.github.io/AudioEditing/">[Project page]</a>&nbsp;|&nbsp;177    <a href="https://github.com/HilaManor/AudioEditingCode">[Code]</a>178</h3>179 180 181<p style="font-size: 0.9rem; margin: 0rem; line-height: 1.2em; margin-top:1em">182For faster inference without waiting in queue, you may duplicate the space and upgrade to GPU in settings.183<a href="https://huggingface.co/spaces/hilamanor/audioEditing?duplicate=true">184<img style="margin-top: 0em; margin-bottom: 0em; display:inline" src="https://bit.ly/3gLdBN6" alt="Duplicate Space" ></a>185</p>186 187"""188 189help = """190<div style="font-size:medium">191<b>Instructions:</b><br>192<ul style="line-height: normal">193<li>You must provide an input audio and a target prompt to edit the audio. </li>194<li>T<sub>start</sub> is used to control the tradeoff between fidelity to the original signal and text-adhearance.195Lower value -> favor fidelity. Higher value -> apply a stronger edit.</li>196<li>Make sure that you use an AudioLDM2 version that is suitable for your input audio.197For example, use the music version for music and the large version for general audio.198</li>199<li>You can additionally provide a source prompt to guide even further the editing process.</li>200<li>Longer input will take more time.</li>201<li><strong>Unlimited length</strong>: This space automatically trims input audio to a maximum length of 30 seconds.202For unlimited length, duplicated the space, and remove the trimming by changing the code.203Specifically, in the <code style="display:inline; background-color: lightgrey; ">load_audio</code> function in the <code style="display:inline; background-color: lightgrey; ">utils.py</code> file,204change <code style="display:inline; background-color: lightgrey; ">duration = min(audioldm.utils.get_duration(audio_path), 30)</code> to 205<code style="display:inline; background-color: lightgrey; ">duration = audioldm.utils.get_duration(audio_path)</code>.206</ul>207</div>208 209"""210 211with gr.Blocks(css='style.css', delete_cache=(3600, 3600)) as demo:212    def reset_do_inversion(do_inversion_user, do_inversion):213        # do_inversion = gr.State(value=True)214        do_inversion = True215        do_inversion_user = True216        return do_inversion_user, do_inversion217 218    # handle the case where the user clicked the button but the inversion was not done219    def clear_do_inversion_user(do_inversion_user):220        do_inversion_user = False221        return do_inversion_user222    def post_match_do_inversion(do_inversion_user, do_inversion):223        if do_inversion_user:224            do_inversion = True225            do_inversion_user = False226        return do_inversion_user, do_inversion227        228    229    gr.HTML(intro)230    # wts = gr.State()231    # zs = gr.State()232    wtszs = gr.State()233    cache_dir = gr.State(demo.GRADIO_CACHE)234    saved_inv_model = gr.State()235    # current_loaded_model = gr.State(value="cvssp/audioldm2-music")236    # ldm_stable = load_model("cvssp/audioldm2-music", device, 200)237    # ldm_stable = gr.State(value=ldm_stable)238    do_inversion = gr.State(value=True)  # To save some runtime when editing the same thing over and over239    do_inversion_user = gr.State(value=False)240 241    with gr.Group():242        gr.Markdown("πŸ’‘ **note**: input longer than **30 sec** is automatically trimmed (for unlimited input, see the Help section below)")243        with gr.Row():244            input_audio = gr.Audio(sources=["upload", "microphone"], type="filepath", editable=True, label="Input Audio",245                                   interactive=True, scale=1)246            output_audio = gr.Audio(label="Edited Audio", interactive=False, scale=1)247 248    with gr.Row():249        tar_prompt = gr.Textbox(label="Prompt", info="Describe your desired edited output",250                                placeholder="a recording of a happy upbeat arcade game soundtrack",251                                lines=2, interactive=True)252 253    with gr.Row():254        t_start = gr.Slider(minimum=15, maximum=85, value=45, step=1, label="T-start (%)", interactive=True, scale=3,255                            info="Lower T-start -> closer to original audio. Higher T-start -> stronger edit.")256        # model_id = gr.Radio(label="AudioLDM2 Version",257        model_id = gr.Dropdown(label="AudioLDM2 Version",258                               choices=["cvssp/audioldm2",259                                        "cvssp/audioldm2-large",260                                        "cvssp/audioldm2-music"],261                               info="Choose a checkpoint suitable for your intended audio and edit",262                               value="cvssp/audioldm2-music", interactive=True, type="value", scale=2)263 264    with gr.Row():265        with gr.Column():266            submit = gr.Button("Edit")267 268    with gr.Accordion("More Options", open=False):269        with gr.Row():270            src_prompt = gr.Textbox(label="Source Prompt", lines=2, interactive=True,271                                    info="Optional: Describe the original audio input",272                                    placeholder="A recording of a happy upbeat classical music piece",)273 274        with gr.Row():275            cfg_scale_src = gr.Number(value=3, minimum=0.5, maximum=25, precision=None,276                                      label="Source Guidance Scale", interactive=True, scale=1)277            cfg_scale_tar = gr.Number(value=12, minimum=0.5, maximum=25, precision=None,278                                      label="Target Guidance Scale", interactive=True, scale=1)279            steps = gr.Number(value=50, step=1, minimum=20, maximum=300,280                              info="Higher values (e.g. 200) yield higher-quality generation.",281                              label="Num Diffusion Steps", interactive=True, scale=1)282        with gr.Row():283            seed = gr.Number(value=0, precision=0, label="Seed", interactive=True)284            randomize_seed = gr.Checkbox(label='Randomize seed', value=False)285            length = gr.Number(label="Length", interactive=False, visible=False)286 287    with gr.Accordion("HelpπŸ’‘", open=False):288        gr.HTML(help)289 290    submit.click(291        fn=randomize_seed_fn,292        inputs=[seed, randomize_seed],293        outputs=[seed], queue=False).then(294            fn=clear_do_inversion_user, inputs=[do_inversion_user], outputs=[do_inversion_user]).then(295           fn=edit,296           inputs=[cache_dir,297                   input_audio,298                   model_id,299                   do_inversion,300                   #    current_loaded_model, ldm_stable,301                   #    wts, zs,302                   wtszs,303                   saved_inv_model,304                   src_prompt,305                   tar_prompt,306                   steps,307                   cfg_scale_src,308                   cfg_scale_tar,309                   t_start,310                   randomize_seed311                   ],312           outputs=[output_audio, wtszs,313                    saved_inv_model, do_inversion]  # , current_loaded_model, ldm_stable],314        ).then(post_match_do_inversion, inputs=[do_inversion_user, do_inversion], outputs=[do_inversion_user, do_inversion]315               ).then(lambda x: (demo.temp_file_sets.append(set([str(gr.utils.abspath(x))])) if type(x) is str else None),316               inputs=wtszs)317 318    # demo.move_resource_to_block_cache(wtszs.value)319 320    # If sources changed we have to rerun inversion321    input_audio.change(fn=reset_do_inversion, inputs=[do_inversion_user, do_inversion], outputs=[do_inversion_user, do_inversion])322    src_prompt.change(fn=reset_do_inversion, inputs=[do_inversion_user, do_inversion], outputs=[do_inversion_user, do_inversion])323    model_id.change(fn=reset_do_inversion, inputs=[do_inversion_user, do_inversion], outputs=[do_inversion_user, do_inversion])324    cfg_scale_src.change(fn=reset_do_inversion, inputs=[do_inversion_user, do_inversion], outputs=[do_inversion_user, do_inversion])325    steps.change(fn=reset_do_inversion, inputs=[do_inversion_user, do_inversion], outputs=[do_inversion_user, do_inversion])326 327    gr.Examples(328        label="Examples",329        examples=get_example(),330        inputs=[input_audio, src_prompt, tar_prompt, t_start, model_id, length, output_audio],331        outputs=[output_audio]332    )333 334    demo.queue()335    demo.launch()336