VenusM1nT/HCIProject
1
1import gradio as gr2from PIL import Image3import time4import threading5import tqdm6# -------------------------------model--------------------------7import subprocess8import sys9 10 11def run_model(txt='wqq'):12 print("start running model")13 ckpt_path = '/home/zhutiantian/code/One-DM/One-DM-ckpt.pt'14 dir_path = '/home/zhutiantian/code/One-DM/Generated/English'15 code_path = '/home/zhutiantian/code/One-DM/test.py'16 17 # 构造 Conda 命令18 conda_path = "/home/zhutiantian/anaconda3/condabin/conda"19 command = f"export LD_LIBRARY_PATH=/usr/lib/wsl/lib:$LD_LIBRARY_PATH && cd /home/zhutiantian/code/One-DM && {conda_path} run -n torch13 CUDA_VISIBLE_DEVICES=0 torchrun --nproc_per_node=1 {code_path} --one_dm {ckpt_path} --generate_type oov_u --dir {dir_path} --input_text {txt}"20 21 subprocess.run(command, shell=True)22 23 24# -------------------------------gradio-------------------------25 26# 所有button响应的函数27 28# 任务处理函数29def process_task(txt, p1, p2, p3, p4, p5):30 # 这里需要添加对txt的检查!31 # txt为空、非英文字母时,返回错误,要求重新填写32 # p3、p5 非空33 # p1: 时间,p2:图片数量,p4:步数34 # p3:等于 "DDIM" 或 "DDPM",生成方式35 # p5:等于 "iv_s","iv_u","oov_s","oov_u",生成类型36 37 38 # 启动一个 15 秒的命令模拟39 def execute_command():40 # 这里用time.sleep(50)来模拟命令执行41 # get_wsl.open_wsl_and_run_model() # 单独修改UI的时候把这里换成time.sleep(15)42 43 # run_model(txt)44 time.sleep(5)45 print("命令执行完毕!")46 47 # 启动执行命令的线程48 command_thread = threading.Thread(target=execute_command)49 command_thread.start()50 time.sleep(5)51 command_thread.join() # 等待命令执行完成52 53 # img_path = f'/home/zhutiantian/code/One-DM/Generated/English/oov_u/168/{txt}.png'54 img_path = ["img/bg.png"]55 56 return img_path, gr.update(visible=True)57 58 59def func(txt, p1, p2, p3, p4, progress=gr.Progress()):60 progress(0, desc="Starting")61 time.sleep(1)62 progress(0.3, desc="Progressing")63 time.sleep(p1)64 progress(1, desc="Completed")65 img = Image.open('img/bg.png')66 67 time.sleep(2)68 # 返回显示评价区域69 return img, gr.update(visible=True)70 71 72def funcgal(txt, p1, p2, p3, p4, progress=gr.Progress()):73 # imggallery = ["img/1.jpg", "img/2.jpg", "img/3.jpg", "img/4.jpg", "img/5.jpg", "img/6.jpg", "img/7.jpg", "img/8.jpg", "img/9.jpg", "img/10.jpg", "img/11.jpg", "img/12.jpg", "img/13.jpg", "img/14.jpg", "img/15.jpg", "img/16.jpg", "img/17.jpg", "img/18.jpg", "img/bg.png", "img/sample.png"]74 imggallery = ["img/1.jpg"]75 76 return imggallery, gr.update(visible=True)77 78 79# 主函数,调用demo80 81def main():82 with gr.Blocks(theme='NoCrypt/miku') as demo:83 feedback_visible = gr.State(False)84 85 def submit_feedback(feedback):86 # 隐藏评价区域,并返回提示信息87 gr.Info("感谢您的评价")88 return gr.update(visible=False)89 90 with gr.Column():91 introtext1 = gr.Markdown( # 在此输入描述,使用 Markdown92 """93 ## “一眼临摹”*手写风格迁移* 94 该项目旨在以手写文本图像为基础,学习其手写字迹风格并生成对应风格的特定文本内容。 95 该项目以论文 [One-Shot Diffusion Mimicker for Handwritten Text Generation](https://arxiv.org/abs/2409.04004) 为基础。该论文详细介绍了如何通过提取单一参考样本的高频信息来改进样式提取,并在此基础上生成对应的手写文本图像。部分代码修改自该论文的 [源代码](https://github.com/dailenson/One-DM)。96 """97 )98 with gr.Group(visible=False) as feedback_group:99 feedback = gr.Textbox(label="请输入您的使用评价")100 submit_btn = gr.Button("提交")101 with gr.Row():102 with gr.Column():103 textinfo = gr.Textbox(label="此处输入你要生成的文字")104 with gr.Row():105 # inputimage = gr.Image(label="此处上传你要生成的字迹风格图像")106 with gr.Column():107 para1 = gr.Slider(label="运行 GPU 卡数", minimum = 1, maximum = 4, step = 1)108 para2 = gr.Slider(label="生成风格数量", minimum = 10, maximum = 150)109 para3 = gr.Radio(["DDIM", "DDPM"], label = "生成方式", info = "解释")110 para4 = gr.Slider(label="生成 sample 步数", minimum = 50, maximum = 1000)111 para5 = gr.Radio(["iv_s", "iv_u", "oov_s", "oov_u"], label="生成类型", info = "解释")112 with gr.Row():113 # genebutton = gr.Button("生成用户字体风格")114 genegallery = gr.Button("生成示例字体风格")115 with gr.Column():116 # outputimage = gr.Image(label="用户字体风格图片")117 outputgallery = gr.Gallery(label="示例字体风格图片", columns=5)118 with gr.Row():119 introtext2 = gr.Image(label="演示图片,来自 One-DM", value="img/intro1.jpg") # 在此修改描述图片路径120 introtext3 = gr.Image(label="演示图片,来自 One-DM", value="img/intro2.png")121 with gr.Row():122 gr.Markdown("""123 <font size=4>124 这里写第一段。125 </font>126 """)127 gr.Markdown("""128 <font size=4>129 这里写第二段。130 </font>131 """)132 133 134 submit_btn.click(135 fn=submit_feedback,136 inputs=[feedback],137 outputs=[feedback_group]138 )139 140 # genebutton.click(141 # fn=process_task,142 # inputs=[textinfo, para1, para2, para3, para4],143 # outputs=[outputimage, feedback_group]144 # )145 146 genegallery.click(147 fn=process_task,148 inputs=[textinfo, para1, para2, para3, para4, para5],149 outputs=[outputgallery, feedback_group]150 )151 152 # demo.launch(server_name="172.30.180.28", server_port=45632, debug=True, show_error=True)153 demo.launch()154 155 156if __name__ == "__main__":157 main()