LobsterQQQ/Text-Image-3D_Model
1
1import os2from PIL import Image3import torch4 5from point_e.diffusion.configs import DIFFUSION_CONFIGS, diffusion_from_config6from point_e.diffusion.sampler import PointCloudSampler7from point_e.models.download import load_checkpoint8from point_e.models.configs import MODEL_CONFIGS, model_from_config9from point_e.util.plotting import plot_point_cloud10from point_e.util.pc_to_mesh import marching_cubes_mesh11 12import skimage.measure13 14from pyntcloud import PyntCloud15import matplotlib.colors16import plotly.graph_objs as go17 18import trimesh19 20import gradio as gr21 22 23state = ""24device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')25 26def set_state(s):27 print(s)28 global state29 state = s30 31def get_state():32 return state33 34set_state('Creating txt2mesh model...')35t2m_name = 'base40M-textvec'36t2m_model = model_from_config(MODEL_CONFIGS[t2m_name], device)37t2m_model.eval()38base_diffusion_t2m = diffusion_from_config(DIFFUSION_CONFIGS[t2m_name])39 40set_state('Downloading txt2mesh checkpoint...')41t2m_model.load_state_dict(load_checkpoint(t2m_name, device))42 43 44def load_img2mesh_model(model_name):45 set_state(f'Creating img2mesh model {model_name}...')46 i2m_name = model_name47 i2m_model = model_from_config(MODEL_CONFIGS[i2m_name], device)48 i2m_model.eval()49 base_diffusion_i2m = diffusion_from_config(DIFFUSION_CONFIGS[i2m_name])50 51 set_state(f'Downloading img2mesh checkpoint {model_name}...')52 i2m_model.load_state_dict(load_checkpoint(i2m_name, device))53 54 return i2m_model, base_diffusion_i2m55 56img2mesh_model_name = 'base40M' #'base300M' #'base1B'57i2m_model, base_diffusion_i2m = load_img2mesh_model(img2mesh_model_name)58 59 60set_state('Creating upsample model...')61upsampler_model = model_from_config(MODEL_CONFIGS['upsample'], device)62upsampler_model.eval()63upsampler_diffusion = diffusion_from_config(DIFFUSION_CONFIGS['upsample'])64 65set_state('Downloading upsampler checkpoint...')66upsampler_model.load_state_dict(load_checkpoint('upsample', device))67 68set_state('Creating SDF model...')69sdf_name = 'sdf'70sdf_model = model_from_config(MODEL_CONFIGS[sdf_name], device)71sdf_model.eval()72 73set_state('Loading SDF model...')74sdf_model.load_state_dict(load_checkpoint(sdf_name, device))75 76stable_diffusion = gr.Blocks.load(name="spaces/runwayml/stable-diffusion-v1-5")77 78 79set_state('')80 81def get_sampler(model_name, txt2obj, guidance_scale):82 83 global img2mesh_model_name84 global base_diffusion_i2m85 global i2m_model86 if model_name != img2mesh_model_name:87 img2mesh_model_name = model_name88 i2m_model, base_diffusion_i2m = load_img2mesh_model(model_name)89 90 return PointCloudSampler(91 device=device,92 models=[t2m_model if txt2obj else i2m_model, upsampler_model],93 diffusions=[base_diffusion_t2m if txt2obj else base_diffusion_i2m, upsampler_diffusion],94 num_points=[1024, 4096 - 1024],95 aux_channels=['R', 'G', 'B'],96 guidance_scale=[guidance_scale, 0.0 if txt2obj else guidance_scale],97 model_kwargs_key_filter=('texts', '') if txt2obj else ("*",)98 )99 100def generate_txt2img(prompt):101 102 prompt = f"“a 3d rendering of {prompt}, full view, white background"103 gallery_dir = stable_diffusion(prompt, fn_index=2)104 imgs = [os.path.join(gallery_dir, img) for img in os.listdir(gallery_dir) if os.path.splitext(img)[1] == '.jpg']105 106 return imgs[0], gr.update(visible=True)107 108def generate_3D(input, model_name='base40M', guidance_scale=3.0, grid_size=32):109 110 set_state('Entered generate function...')111 112 if isinstance(input, Image.Image):113 input = prepare_img(input)114 115 # if input is a string, it's a text prompt116 sampler = get_sampler(model_name, txt2obj=True if isinstance(input, str) else False, guidance_scale=guidance_scale)117 118 # Produce a sample from the model.119 set_state('Sampling...')120 samples = None121 kw_args = dict(texts=[input]) if isinstance(input, str) else dict(images=[input])122 for x in sampler.sample_batch_progressive(batch_size=1, model_kwargs=kw_args):123 samples = x124 125 set_state('Converting to point cloud...')126 pc = sampler.output_to_point_clouds(samples)[0]127 128 set_state('Saving point cloud...')129 with open("point_cloud.ply", "wb") as f:130 pc.write_ply(f)131 132 set_state('Converting to mesh...')133 save_ply(pc, 'mesh.ply', grid_size)134 135 set_state('')136 137 return pc_to_plot(pc), ply_to_obj('mesh.ply', '3d_model.obj'), gr.update(value=['3d_model.obj', 'mesh.ply', 'point_cloud.ply'], visible=True)138 139def prepare_img(img):140 141 w, h = img.size142 if w > h:143 img = img.crop((w - h) / 2, 0, w - (w - h) / 2, h)144 else:145 img = img.crop((0, (h - w) / 2, w, h - (h - w) / 2))146 147 # resize to 256x256148 img = img.resize((256, 256))149 150 return img151 152def pc_to_plot(pc):153 154 return go.Figure(155 data=[156 go.Scatter3d(157 x=pc.coords[:,0], y=pc.coords[:,1], z=pc.coords[:,2], 158 mode='markers',159 marker=dict(160 size=2,161 color=['rgb({},{},{})'.format(r,g,b) for r,g,b in zip(pc.channels["R"], pc.channels["G"], pc.channels["B"])],162 )163 )164 ],165 layout=dict(166 scene=dict(xaxis=dict(visible=False), yaxis=dict(visible=False), zaxis=dict(visible=False))167 ),168 )169 170def ply_to_obj(ply_file, obj_file):171 mesh = trimesh.load(ply_file)172 mesh.export(obj_file)173 174 return obj_file175 176def save_ply(pc, file_name, grid_size):177 178 # Produce a mesh (with vertex colors)179 mesh = marching_cubes_mesh(180 pc=pc,181 model=sdf_model,182 batch_size=4096,183 grid_size=grid_size, # increase to 128 for resolution used in evals184 fill_vertex_channels=True,185 progress=True,186 )187 188 # Write the mesh to a PLY file to import into some other program.189 with open(file_name, 'wb') as f:190 mesh.write_ply(f)191 192 193with gr.Blocks() as app:194 195 196 with gr.Row():197 with gr.Column():198 with gr.Tab("Text to 3D"):199 prompt = gr.Textbox(label="Prompt", placeholder="A cactus in a pot")200 btn_generate_txt2obj = gr.Button(value="Generate")201 202 with gr.Tab("Image to 3D"):203 img = gr.Image(label="Image")204 gr.Markdown("Best results with images of 3D objects with no shadows on a white background.")205 btn_generate_img2obj = gr.Button(value="Generate")206 207 with gr.Tab("Text to Image to 3D"):208 gr.Markdown("Generate an image with Stable Diffusion, then convert it to 3D. Just enter the object you want to generate.")209 prompt_sd = gr.Textbox(label="Prompt", placeholder="a 3d rendering of [your prompt], full view, white background")210 btn_generate_txt2sd = gr.Button(value="Generate image")211 img_sd = gr.Image(label="Image")212 btn_generate_sd2obj = gr.Button(value="Convert to 3D", visible=False)213 214 with gr.Accordion("Advanced settings", open=False):215 dropdown_models = gr.Dropdown(label="Model", value="base40M", choices=["base40M", "base300M"]) #, "base1B"])216 guidance_scale = gr.Slider(label="Guidance scale", value=3.0, minimum=3.0, maximum=10.0, step=0.1)217 grid_size = gr.Slider(label="Grid size (for .obj 3D model)", value=32, minimum=16, maximum=128, step=16)218 219 with gr.Column():220 plot = gr.Plot(label="Point cloud")221 # btn_pc_to_obj = gr.Button(value="Convert to OBJ", visible=False)222 model_3d = gr.Model3D(value=None)223 file_out = gr.File(label="Files", visible=False)224 225 # state_info = state_info = gr.Textbox(label="State", show_label=False).style(container=False)226 227 228 # inputs = [dropdown_models, prompt, img, guidance_scale, grid_size]229 outputs = [plot, model_3d, file_out]230 231 prompt.submit(generate_3D, inputs=[prompt, dropdown_models, guidance_scale, grid_size], outputs=outputs)232 btn_generate_txt2obj.click(generate_3D, inputs=[prompt, dropdown_models, guidance_scale, grid_size], outputs=outputs)233 234 btn_generate_img2obj.click(generate_3D, inputs=[img, dropdown_models, guidance_scale, grid_size], outputs=outputs)235 236 prompt_sd.submit(generate_txt2img, inputs=prompt_sd, outputs=[img_sd, btn_generate_sd2obj])237 btn_generate_txt2sd.click(generate_txt2img, inputs=prompt_sd, outputs=[img_sd, btn_generate_sd2obj], queue=False)238 btn_generate_sd2obj.click(generate_3D, inputs=[img, dropdown_models, guidance_scale, grid_size], outputs=outputs)239 240 # btn_pc_to_obj.click(ply_to_obj, inputs=plot, outputs=[model_3d, file_out])241 242 243 # app.load(get_state, inputs=[], outputs=state_info, every=0.5, show_progress=False)244 245 246app.queue(max_size=250, concurrency_count=6).launch()