mirror of
https://github.com/SWivid/F5-TTS.git
synced 2026-07-28 21:35:22 -07:00
v1.0.5 update finetune_gradio.py for clearer guidance
This commit is contained in:
+1
-1
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "f5-tts"
|
||||
version = "1.0.4"
|
||||
version = "1.0.5"
|
||||
description = "F5-TTS: A Fairytaler that Fakes Fluent and Faithful Speech with Flow Matching"
|
||||
readme = "README.md"
|
||||
license = {text = "MIT License"}
|
||||
|
||||
@@ -758,9 +758,9 @@ This is {"a local web UI for [F5 TTS](https://github.com/SWivid/F5-TTS)" if not
|
||||
|
||||
The checkpoints currently support English and Chinese.
|
||||
|
||||
If you're having issues, try converting your reference audio to WAV or MP3, clipping it to 15s with ✂ in the bottom right corner (otherwise might have non-optimal auto-trimmed result).
|
||||
If you're having issues, try converting your reference audio to WAV or MP3, clipping it to 12s with ✂ in the bottom right corner (otherwise might have non-optimal auto-trimmed result).
|
||||
|
||||
**NOTE: Reference text will be automatically transcribed with Whisper if not provided. For best results, keep your reference clips short (<15s). Ensure the audio is fully uploaded before generating.**
|
||||
**NOTE: Reference text will be automatically transcribed with Whisper if not provided. For best results, keep your reference clips short (<12s). Ensure the audio is fully uploaded before generating.**
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
@@ -302,7 +302,7 @@ def preprocess_ref_audio_text(ref_audio_orig, ref_text, clip_short=True, show_in
|
||||
non_silent_wave = AudioSegment.silent(duration=0)
|
||||
for non_silent_seg in non_silent_segs:
|
||||
if len(non_silent_wave) > 6000 and len(non_silent_wave + non_silent_seg) > 12000:
|
||||
show_info("Audio is over 15s, clipping short. (1)")
|
||||
show_info("Audio is over 12s, clipping short. (1)")
|
||||
break
|
||||
non_silent_wave += non_silent_seg
|
||||
|
||||
@@ -314,7 +314,7 @@ def preprocess_ref_audio_text(ref_audio_orig, ref_text, clip_short=True, show_in
|
||||
non_silent_wave = AudioSegment.silent(duration=0)
|
||||
for non_silent_seg in non_silent_segs:
|
||||
if len(non_silent_wave) > 6000 and len(non_silent_wave + non_silent_seg) > 12000:
|
||||
show_info("Audio is over 15s, clipping short. (2)")
|
||||
show_info("Audio is over 12s, clipping short. (2)")
|
||||
break
|
||||
non_silent_wave += non_silent_seg
|
||||
|
||||
@@ -323,7 +323,7 @@ def preprocess_ref_audio_text(ref_audio_orig, ref_text, clip_short=True, show_in
|
||||
# 3. if no proper silence found for clipping
|
||||
if len(aseg) > 12000:
|
||||
aseg = aseg[:12000]
|
||||
show_info("Audio is over 15s, clipping short. (3)")
|
||||
show_info("Audio is over 12s, clipping short. (3)")
|
||||
|
||||
aseg = remove_silence_edges(aseg) + AudioSegment.silent(duration=50)
|
||||
aseg.export(f.name, format="wav")
|
||||
|
||||
@@ -40,15 +40,15 @@ def parse_args():
|
||||
parser.add_argument("--grad_accumulation_steps", type=int, default=1, help="Gradient accumulation steps")
|
||||
parser.add_argument("--max_grad_norm", type=float, default=1.0, help="Max gradient norm for clipping")
|
||||
parser.add_argument("--epochs", type=int, default=100, help="Number of training epochs")
|
||||
parser.add_argument("--num_warmup_updates", type=int, default=300, help="Warmup updates")
|
||||
parser.add_argument("--save_per_updates", type=int, default=10000, help="Save checkpoint every X updates")
|
||||
parser.add_argument("--num_warmup_updates", type=int, default=20000, help="Warmup updates")
|
||||
parser.add_argument("--save_per_updates", type=int, default=50000, help="Save checkpoint every N updates")
|
||||
parser.add_argument(
|
||||
"--keep_last_n_checkpoints",
|
||||
type=int,
|
||||
default=-1,
|
||||
help="-1 to keep all, 0 to not save intermediate, > 0 to keep last N checkpoints",
|
||||
)
|
||||
parser.add_argument("--last_per_updates", type=int, default=50000, help="Save last checkpoint every X updates")
|
||||
parser.add_argument("--last_per_updates", type=int, default=5000, help="Save last checkpoint every N updates")
|
||||
parser.add_argument("--finetune", action="store_true", help="Use Finetune")
|
||||
parser.add_argument("--pretrain", type=str, default=None, help="the path to the checkpoint")
|
||||
parser.add_argument(
|
||||
@@ -65,7 +65,7 @@ def parse_args():
|
||||
action="store_true",
|
||||
help="Log inferenced samples per ckpt save updates",
|
||||
)
|
||||
parser.add_argument("--logger", type=str, default=None, choices=["wandb", "tensorboard"], help="logger")
|
||||
parser.add_argument("--logger", type=str, default=None, choices=[None, "wandb", "tensorboard"], help="logger")
|
||||
parser.add_argument(
|
||||
"--bnb_optimizer",
|
||||
action="store_true",
|
||||
|
||||
@@ -120,11 +120,11 @@ def load_settings(project_name):
|
||||
default_settings = {
|
||||
"exp_name": "F5TTS_v1_Base",
|
||||
"learning_rate": 1e-5,
|
||||
"batch_size_per_gpu": 1,
|
||||
"batch_size_type": "sample",
|
||||
"batch_size_per_gpu": 3200,
|
||||
"batch_size_type": "frame",
|
||||
"max_samples": 64,
|
||||
"grad_accumulation_steps": 4,
|
||||
"max_grad_norm": 1,
|
||||
"grad_accumulation_steps": 1,
|
||||
"max_grad_norm": 1.0,
|
||||
"epochs": 100,
|
||||
"num_warmup_updates": 100,
|
||||
"save_per_updates": 500,
|
||||
@@ -134,8 +134,8 @@ def load_settings(project_name):
|
||||
"file_checkpoint_train": "",
|
||||
"tokenizer_type": "pinyin",
|
||||
"tokenizer_file": "",
|
||||
"mixed_precision": "none",
|
||||
"logger": "wandb",
|
||||
"mixed_precision": "fp16",
|
||||
"logger": "none",
|
||||
"bnb_optimizer": False,
|
||||
}
|
||||
|
||||
@@ -361,27 +361,27 @@ def terminate_process(pid):
|
||||
|
||||
|
||||
def start_training(
|
||||
dataset_name="",
|
||||
exp_name="F5TTS_v1_Base",
|
||||
learning_rate=1e-5,
|
||||
batch_size_per_gpu=1,
|
||||
batch_size_type="sample",
|
||||
max_samples=64,
|
||||
grad_accumulation_steps=4,
|
||||
max_grad_norm=1.0,
|
||||
epochs=100,
|
||||
num_warmup_updates=100,
|
||||
save_per_updates=500,
|
||||
keep_last_n_checkpoints=-1,
|
||||
last_per_updates=100,
|
||||
finetune=True,
|
||||
file_checkpoint_train="",
|
||||
tokenizer_type="pinyin",
|
||||
tokenizer_file="",
|
||||
mixed_precision="fp16",
|
||||
stream=False,
|
||||
logger="wandb",
|
||||
ch_8bit_adam=False,
|
||||
dataset_name,
|
||||
exp_name,
|
||||
learning_rate,
|
||||
batch_size_per_gpu,
|
||||
batch_size_type,
|
||||
max_samples,
|
||||
grad_accumulation_steps,
|
||||
max_grad_norm,
|
||||
epochs,
|
||||
num_warmup_updates,
|
||||
save_per_updates,
|
||||
keep_last_n_checkpoints,
|
||||
last_per_updates,
|
||||
finetune,
|
||||
file_checkpoint_train,
|
||||
tokenizer_type,
|
||||
tokenizer_file,
|
||||
mixed_precision,
|
||||
stream,
|
||||
logger,
|
||||
ch_8bit_adam,
|
||||
):
|
||||
global training_process, tts_api, stop_signal
|
||||
|
||||
@@ -458,7 +458,10 @@ def start_training(
|
||||
|
||||
cmd += f" --tokenizer {tokenizer_type}"
|
||||
|
||||
cmd += f" --log_samples --logger {logger}"
|
||||
if logger != "none":
|
||||
cmd += f" --logger {logger}"
|
||||
|
||||
cmd += " --log_samples"
|
||||
|
||||
if ch_8bit_adam:
|
||||
cmd += " --bnb_optimizer"
|
||||
@@ -515,7 +518,7 @@ def start_training(
|
||||
training_process = subprocess.Popen(
|
||||
cmd, shell=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, bufsize=1, env=env
|
||||
)
|
||||
yield "Training started...", gr.update(interactive=False), gr.update(interactive=True)
|
||||
yield "Training started ...", gr.update(interactive=False), gr.update(interactive=True)
|
||||
|
||||
stdout_queue = queue.Queue()
|
||||
stderr_queue = queue.Queue()
|
||||
@@ -584,7 +587,11 @@ def start_training(
|
||||
gr.update(interactive=True),
|
||||
)
|
||||
else:
|
||||
yield "Training complete!", gr.update(interactive=False), gr.update(interactive=True)
|
||||
yield (
|
||||
"Training complete or paused ...",
|
||||
gr.update(interactive=False),
|
||||
gr.update(interactive=True),
|
||||
)
|
||||
break
|
||||
|
||||
# Small sleep to prevent CPU thrashing
|
||||
@@ -598,9 +605,9 @@ def start_training(
|
||||
time.sleep(1)
|
||||
|
||||
if training_process is None:
|
||||
text_info = "train stop"
|
||||
text_info = "Train stopped !"
|
||||
else:
|
||||
text_info = "train complete !"
|
||||
text_info = "Train complete at end !"
|
||||
|
||||
except Exception as e: # Catch all exceptions
|
||||
# Ensure that we reset the training process variable in case of an error
|
||||
@@ -615,11 +622,11 @@ def stop_training():
|
||||
global training_process, stop_signal
|
||||
|
||||
if training_process is None:
|
||||
return "Train not run !", gr.update(interactive=True), gr.update(interactive=False)
|
||||
return "Train not running !", gr.update(interactive=True), gr.update(interactive=False)
|
||||
terminate_process_tree(training_process.pid)
|
||||
# training_process = None
|
||||
stop_signal = True
|
||||
return "train stop", gr.update(interactive=True), gr.update(interactive=False)
|
||||
return "Train stopped !", gr.update(interactive=True), gr.update(interactive=False)
|
||||
|
||||
|
||||
def get_list_projects():
|
||||
@@ -1128,7 +1135,7 @@ def vocab_check(project_name):
|
||||
info = "You can train using your language !"
|
||||
else:
|
||||
vocab_miss = ",".join(miss_symbols)
|
||||
info = f"The following symbols are missing in your language {len(miss_symbols)}\n\n"
|
||||
info = f"The following {len(miss_symbols)} symbols are missing in your language\n\n"
|
||||
|
||||
return info, vocab_miss
|
||||
|
||||
@@ -1215,6 +1222,9 @@ def infer(
|
||||
|
||||
print("update >> ", device_test, file_checkpoint, use_ema)
|
||||
|
||||
if seed == -1: # -1 used for random
|
||||
seed = None
|
||||
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as f:
|
||||
tts_api.infer(
|
||||
ref_file=ref_audio,
|
||||
@@ -1433,9 +1443,9 @@ Skip this step if you have your dataset, metadata.csv, and a folder wavs with al
|
||||
)
|
||||
|
||||
audio_speaker = gr.File(label="Voice", type="filepath", file_count="multiple")
|
||||
txt_lang = gr.Text(label="Language", value="English")
|
||||
txt_lang = gr.Textbox(label="Language", value="English")
|
||||
bt_transcribe = bt_create = gr.Button("Transcribe")
|
||||
txt_info_transcribe = gr.Text(label="Info", value="")
|
||||
txt_info_transcribe = gr.Textbox(label="Info", value="")
|
||||
bt_transcribe.click(
|
||||
fn=transcribe_all,
|
||||
inputs=[cm_project, audio_speaker, txt_lang, ch_manual],
|
||||
@@ -1446,7 +1456,7 @@ Skip this step if you have your dataset, metadata.csv, and a folder wavs with al
|
||||
random_sample_transcribe = gr.Button("Random Sample")
|
||||
|
||||
with gr.Row():
|
||||
random_text_transcribe = gr.Text(label="Text")
|
||||
random_text_transcribe = gr.Textbox(label="Text")
|
||||
random_audio_transcribe = gr.Audio(label="Audio", type="filepath")
|
||||
|
||||
random_sample_transcribe.click(
|
||||
@@ -1461,7 +1471,7 @@ Check the vocabulary for fine-tuning Emilia_ZH_EN to ensure all symbols are incl
|
||||
```""")
|
||||
|
||||
check_button = gr.Button("Check Vocab")
|
||||
txt_info_check = gr.Text(label="Info", value="")
|
||||
txt_info_check = gr.Textbox(label="Info", value="")
|
||||
|
||||
gr.Markdown("""```plaintext
|
||||
Using the extended model, you can finetune to a new language that is missing symbols in the vocab. This creates a new model with a new vocabulary size and saves it in your ckpts/project folder.
|
||||
@@ -1481,7 +1491,7 @@ Using the extended model, you can finetune to a new language that is missing sym
|
||||
txt_count_symbol = gr.Textbox(label="New Vocab Size", value="", scale=1)
|
||||
|
||||
extend_button = gr.Button("Extend")
|
||||
txt_info_extend = gr.Text(label="Info", value="")
|
||||
txt_info_extend = gr.Textbox(label="Info", value="")
|
||||
|
||||
txt_extend.change(vocab_count, inputs=[txt_extend], outputs=[txt_count_symbol])
|
||||
check_button.click(fn=vocab_check, inputs=[cm_project], outputs=[txt_info_check, txt_extend])
|
||||
@@ -1521,8 +1531,8 @@ Skip this step if you have your dataset, raw.arrow, duration.json, and vocab.txt
|
||||
ch_tokenizern = gr.Checkbox(label="Create Vocabulary", value=False, visible=False)
|
||||
|
||||
bt_prepare = bt_create = gr.Button("Prepare")
|
||||
txt_info_prepare = gr.Text(label="Info", value="")
|
||||
txt_vocab_prepare = gr.Text(label="Vocab", value="")
|
||||
txt_info_prepare = gr.Textbox(label="Info", value="")
|
||||
txt_vocab_prepare = gr.Textbox(label="Vocab", value="")
|
||||
|
||||
bt_prepare.click(
|
||||
fn=create_metadata, inputs=[cm_project, ch_tokenizern], outputs=[txt_info_prepare, txt_vocab_prepare]
|
||||
@@ -1531,7 +1541,7 @@ Skip this step if you have your dataset, raw.arrow, duration.json, and vocab.txt
|
||||
random_sample_prepare = gr.Button("Random Sample")
|
||||
|
||||
with gr.Row():
|
||||
random_text_prepare = gr.Text(label="Tokenizer")
|
||||
random_text_prepare = gr.Textbox(label="Tokenizer")
|
||||
random_audio_prepare = gr.Audio(label="Audio", type="filepath")
|
||||
|
||||
random_sample_prepare.click(
|
||||
@@ -1544,50 +1554,60 @@ The auto-setting is still experimental. Set a large value of epoch if not sure;
|
||||
If you encounter a memory error, try reducing the batch size per GPU to a smaller number.
|
||||
```""")
|
||||
with gr.Row():
|
||||
bt_calculate = bt_create = gr.Button("Auto Settings")
|
||||
exp_name = gr.Radio(label="Model", choices=["F5TTS_v1_Base", "F5TTS_Base", "E2TTS_Base"])
|
||||
tokenizer_file = gr.Textbox(label="Tokenizer File")
|
||||
file_checkpoint_train = gr.Textbox(label="Path to the Pretrained Checkpoint")
|
||||
|
||||
with gr.Row():
|
||||
ch_finetune = bt_create = gr.Checkbox(label="Finetune")
|
||||
lb_samples = gr.Label(label="Samples")
|
||||
batch_size_type = gr.Radio(label="Batch Size Type", choices=["frame", "sample"], value="frame")
|
||||
bt_calculate = bt_create = gr.Button("Auto Settings")
|
||||
|
||||
with gr.Row():
|
||||
ch_finetune = bt_create = gr.Checkbox(label="Finetune", value=True)
|
||||
tokenizer_file = gr.Textbox(label="Tokenizer File", value="")
|
||||
file_checkpoint_train = gr.Textbox(label="Path to the Pretrained Checkpoint", value="")
|
||||
epochs = gr.Number(label="Epochs")
|
||||
learning_rate = gr.Number(label="Learning Rate", step=0.5e-5)
|
||||
max_grad_norm = gr.Number(label="Max Gradient Norm")
|
||||
num_warmup_updates = gr.Number(label="Warmup Updates")
|
||||
|
||||
with gr.Row():
|
||||
exp_name = gr.Radio(
|
||||
label="Model", choices=["F5TTS_v1_Base", "F5TTS_Base", "E2TTS_Base"], value="F5TTS_v1_Base"
|
||||
batch_size_type = gr.Radio(
|
||||
label="Batch Size Type",
|
||||
choices=["frame", "sample"],
|
||||
info="frame is calculated as seconds * sampling_rate / hop_length",
|
||||
)
|
||||
learning_rate = gr.Number(label="Learning Rate", value=1e-5, step=1e-5)
|
||||
batch_size_per_gpu = gr.Number(label="Batch Size per GPU", info="N frames or N samples")
|
||||
grad_accumulation_steps = gr.Number(
|
||||
label="Gradient Accumulation Steps", info="Effective batch size is multiplied by this value"
|
||||
)
|
||||
max_samples = gr.Number(label="Max Samples", info="Maximum number of samples per single GPU batch")
|
||||
|
||||
with gr.Row():
|
||||
batch_size_per_gpu = gr.Number(label="Batch Size per GPU", value=3200)
|
||||
max_samples = gr.Number(label="Max Samples", value=64)
|
||||
|
||||
with gr.Row():
|
||||
grad_accumulation_steps = gr.Number(label="Gradient Accumulation Steps", value=1)
|
||||
max_grad_norm = gr.Number(label="Max Gradient Norm", value=1.0)
|
||||
|
||||
with gr.Row():
|
||||
epochs = gr.Number(label="Epochs", value=100)
|
||||
num_warmup_updates = gr.Number(label="Warmup Updates", value=100)
|
||||
|
||||
with gr.Row():
|
||||
save_per_updates = gr.Number(label="Save per Updates", value=500)
|
||||
save_per_updates = gr.Number(
|
||||
label="Save per Updates",
|
||||
info="Save intermediate checkpoints every N updates",
|
||||
minimum=10,
|
||||
)
|
||||
keep_last_n_checkpoints = gr.Number(
|
||||
label="Keep Last N Checkpoints",
|
||||
value=-1,
|
||||
step=1,
|
||||
precision=0,
|
||||
info="-1 to keep all, 0 to not save intermediate, > 0 to keep last N checkpoints",
|
||||
info="-1 to keep all, 0 to not save intermediate, > 0 to keep last N",
|
||||
minimum=-1,
|
||||
)
|
||||
last_per_updates = gr.Number(label="Last per Updates", value=100)
|
||||
last_per_updates = gr.Number(
|
||||
label="Last per Updates",
|
||||
info="Save latest checkpoint with suffix _last.pt every N updates",
|
||||
minimum=10,
|
||||
)
|
||||
gr.Radio(label="") # placeholder
|
||||
|
||||
with gr.Row():
|
||||
ch_8bit_adam = gr.Checkbox(label="Use 8-bit Adam optimizer")
|
||||
mixed_precision = gr.Radio(label="mixed_precision", choices=["none", "fp16", "bf16"], value="fp16")
|
||||
cd_logger = gr.Radio(label="logger", choices=["wandb", "tensorboard"], value="wandb")
|
||||
start_button = gr.Button("Start Training")
|
||||
stop_button = gr.Button("Stop Training", interactive=False)
|
||||
mixed_precision = gr.Radio(label="Mixed Precision", choices=["none", "fp16", "bf16"])
|
||||
cd_logger = gr.Radio(label="Logger", choices=["none", "wandb", "tensorboard"])
|
||||
with gr.Column():
|
||||
start_button = gr.Button("Start Training")
|
||||
stop_button = gr.Button("Stop Training", interactive=False)
|
||||
|
||||
if projects_selelect is not None:
|
||||
(
|
||||
@@ -1634,7 +1654,7 @@ If you encounter a memory error, try reducing the batch size per GPU to a smalle
|
||||
ch_8bit_adam.value = bnb_optimizer_value
|
||||
|
||||
ch_stream = gr.Checkbox(label="Stream Output Experiment", value=True)
|
||||
txt_info_train = gr.Text(label="Info", value="")
|
||||
txt_info_train = gr.Textbox(label="Info", value="")
|
||||
|
||||
list_audios, select_audio = get_audio_project(projects_selelect, False)
|
||||
|
||||
@@ -1763,7 +1783,7 @@ If you encounter a memory error, try reducing the batch size per GPU to a smalle
|
||||
|
||||
with gr.TabItem("Test Model"):
|
||||
gr.Markdown("""```plaintext
|
||||
SOS: Check the use_ema setting (True or False) for your model to see what works best for you. use seed -1 from random
|
||||
Check the use_ema setting (True or False) for your model to see what works best for you. Set seed to -1 for random.
|
||||
```""")
|
||||
exp_name = gr.Radio(
|
||||
label="Model", choices=["F5TTS_v1_Base", "F5TTS_Base", "E2TTS_Base"], value="F5TTS_v1_Base"
|
||||
@@ -1773,11 +1793,13 @@ SOS: Check the use_ema setting (True or False) for your model to see what works
|
||||
with gr.Row():
|
||||
nfe_step = gr.Number(label="NFE Step", value=32)
|
||||
speed = gr.Slider(label="Speed", value=1.0, minimum=0.3, maximum=2.0, step=0.1)
|
||||
seed = gr.Number(label="Seed", value=-1, minimum=-1)
|
||||
seed = gr.Number(label="Random Seed", value=-1, minimum=-1)
|
||||
remove_silence = gr.Checkbox(label="Remove Silence")
|
||||
|
||||
ch_use_ema = gr.Checkbox(label="Use EMA", value=True)
|
||||
with gr.Row():
|
||||
ch_use_ema = gr.Checkbox(
|
||||
label="Use EMA", value=True, info="Turn off at early stage might offer better results"
|
||||
)
|
||||
cm_checkpoint = gr.Dropdown(
|
||||
choices=list_checkpoints, value=checkpoint_select, label="Checkpoints", allow_custom_value=True
|
||||
)
|
||||
@@ -1785,20 +1807,20 @@ SOS: Check the use_ema setting (True or False) for your model to see what works
|
||||
|
||||
random_sample_infer = gr.Button("Random Sample")
|
||||
|
||||
ref_text = gr.Textbox(label="Ref Text")
|
||||
ref_audio = gr.Audio(label="Audio Ref", type="filepath")
|
||||
gen_text = gr.Textbox(label="Gen Text")
|
||||
ref_text = gr.Textbox(label="Reference Text")
|
||||
ref_audio = gr.Audio(label="Reference Audio", type="filepath")
|
||||
gen_text = gr.Textbox(label="Text to Generate")
|
||||
|
||||
random_sample_infer.click(
|
||||
fn=get_random_sample_infer, inputs=[cm_project], outputs=[ref_text, gen_text, ref_audio]
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
txt_info_gpu = gr.Textbox("", label="Device")
|
||||
seed_info = gr.Text(label="Seed :")
|
||||
check_button_infer = gr.Button("Infer")
|
||||
txt_info_gpu = gr.Textbox("", label="Inference on Device :")
|
||||
seed_info = gr.Textbox(label="Used Random Seed :")
|
||||
check_button_infer = gr.Button("Inference")
|
||||
|
||||
gen_audio = gr.Audio(label="Audio Gen", type="filepath")
|
||||
gen_audio = gr.Audio(label="Generated Audio", type="filepath")
|
||||
|
||||
check_button_infer.click(
|
||||
fn=infer,
|
||||
@@ -1825,10 +1847,10 @@ SOS: Check the use_ema setting (True or False) for your model to see what works
|
||||
gr.Markdown("""```plaintext
|
||||
Reduce the Base model size from 5GB to 1.3GB. The new checkpoint file prunes out optimizer and etc., can be used for inference or finetuning afterward, but not able to resume pretraining.
|
||||
```""")
|
||||
txt_path_checkpoint = gr.Text(label="Path to Checkpoint:")
|
||||
txt_path_checkpoint_small = gr.Text(label="Path to Output:")
|
||||
txt_path_checkpoint = gr.Textbox(label="Path to Checkpoint:")
|
||||
txt_path_checkpoint_small = gr.Textbox(label="Path to Output:")
|
||||
ch_safetensors = gr.Checkbox(label="Safetensors", value="")
|
||||
txt_info_reduse = gr.Text(label="Info", value="")
|
||||
txt_info_reduse = gr.Textbox(label="Info", value="")
|
||||
reduse_button = gr.Button("Reduce")
|
||||
reduse_button.click(
|
||||
fn=extract_and_save_ema_model,
|
||||
|
||||
Reference in New Issue
Block a user