mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
Added Dockerfile and refined the code
This commit is contained in:
@@ -0,0 +1,51 @@
|
||||
FROM nvidia/cuda:11.8.0-devel-ubuntu20.04
|
||||
|
||||
ARG DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV TORCH_CUDA_ARCH_LIST="6.0 6.1 7.0 7.5 8.0 8.6"
|
||||
ENV TCNN_CUDA_ARCHITECTURES=86;80;75;70;61;60
|
||||
ENV FORCE_CUDA=1
|
||||
ENV CUDA_HOME=/usr/local/cuda
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:${LD_LIBRARY_PATH}
|
||||
ENV LIBRARY_PATH=${CUDA_HOME}/lib64/stubs:${LIBRARY_PATH}
|
||||
|
||||
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
curl \
|
||||
git \
|
||||
libegl1-mesa-dev \
|
||||
libgl1-mesa-dev \
|
||||
libgles2-mesa-dev \
|
||||
libglib2.0-0 \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender1 \
|
||||
python-is-python3 \
|
||||
python3-dev \
|
||||
python3-pip \
|
||||
wget \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
ENV HOME=/root
|
||||
ENV PATH=/root/.local/bin:$PATH
|
||||
ENV PYTHONPATH=$HOME/app
|
||||
ENV GRADIO_ALLOW_FLAGGING=never
|
||||
ENV GRADIO_NUM_PORTS=1
|
||||
ENV GRADIO_SERVER_NAME=0.0.0.0
|
||||
ENV GRADIO_THEME=huggingface
|
||||
ENV SYSTEM=spaces
|
||||
|
||||
RUN pip install --upgrade pip setuptools ninja
|
||||
RUN pip install torch==2.2.1 --extra-index-url https://download.pytorch.org/whl/cu118
|
||||
|
||||
RUN python -c "import torch; print(torch.version.cuda)"
|
||||
COPY requirements.txt /tmp/
|
||||
RUN cd /tmp && pip install -r requirements.txt
|
||||
|
||||
# Set the working directory
|
||||
WORKDIR /root/app
|
||||
|
||||
# Copy the current directory contents into the container at /root/app
|
||||
COPY . /root/app
|
||||
@@ -13,6 +13,9 @@ from functools import partial
|
||||
from tsr.system import TSR
|
||||
from tsr.utils import remove_background, resize_foreground, to_gradio_3d_orientation
|
||||
|
||||
import argparse
|
||||
|
||||
|
||||
if torch.cuda.is_available():
|
||||
device = "cuda:0"
|
||||
else:
|
||||
@@ -55,22 +58,25 @@ def preprocess(input_image, do_remove_background, foreground_ratio):
|
||||
return image
|
||||
|
||||
|
||||
def generate(image):
|
||||
def generate(image, mc_resolution, formats=["obj", "glb"]):
|
||||
scene_codes = model(image, device=device)
|
||||
mesh = model.extract_mesh(scene_codes)[0]
|
||||
mesh = model.extract_mesh(scene_codes, resolution=mc_resolution)[0]
|
||||
mesh = to_gradio_3d_orientation(mesh)
|
||||
mesh_path = tempfile.NamedTemporaryFile(suffix=".obj", delete=False)
|
||||
mesh.export(mesh_path.name)
|
||||
return mesh_path.name
|
||||
rv = []
|
||||
for format in formats:
|
||||
mesh_path = tempfile.NamedTemporaryFile(suffix=f".{format}", delete=False)
|
||||
mesh.export(mesh_path.name)
|
||||
rv.append(mesh_path.name)
|
||||
return rv
|
||||
|
||||
|
||||
def run_example(image_pil):
|
||||
preprocessed = preprocess(image_pil, False, 0.9)
|
||||
mesh_name = generate(preprocessed)
|
||||
return preprocessed, mesh_name
|
||||
mesh_name_obj, mesh_name_glb = generate(preprocessed, 256, ["obj", "glb"])
|
||||
return preprocessed, mesh_name_obj, mesh_name_glb
|
||||
|
||||
|
||||
with gr.Blocks() as demo:
|
||||
with gr.Blocks(title="TripoSR") as interface:
|
||||
gr.Markdown(
|
||||
"""
|
||||
# TripoSR Demo
|
||||
@@ -105,14 +111,28 @@ with gr.Blocks() as demo:
|
||||
value=0.85,
|
||||
step=0.05,
|
||||
)
|
||||
mc_resolution = gr.Slider(
|
||||
label="Marching Cubes Resolution",
|
||||
minimum=32,
|
||||
maximum=320,
|
||||
value=256,
|
||||
step=32
|
||||
)
|
||||
with gr.Row():
|
||||
submit = gr.Button("Generate", elem_id="generate", variant="primary")
|
||||
with gr.Column():
|
||||
with gr.Tab("Model"):
|
||||
output_model = gr.Model3D(
|
||||
label="Output Model",
|
||||
with gr.Tab("OBJ"):
|
||||
output_model_obj = gr.Model3D(
|
||||
label="Output Model (OBJ Format)",
|
||||
interactive=False,
|
||||
)
|
||||
gr.Markdown("Note: The model shown here is flipped. Download to get correct results.")
|
||||
with gr.Tab("GLB"):
|
||||
output_model_glb = gr.Model3D(
|
||||
label="Output Model (GLB Format)",
|
||||
interactive=False,
|
||||
)
|
||||
gr.Markdown("Note: The model shown here has a darker appearance. Download to get correct results.")
|
||||
with gr.Row(variant="panel"):
|
||||
gr.Examples(
|
||||
examples=[
|
||||
@@ -131,7 +151,7 @@ with gr.Blocks() as demo:
|
||||
"examples/captured_p.png",
|
||||
],
|
||||
inputs=[input_image],
|
||||
outputs=[processed_image, output_model],
|
||||
outputs=[processed_image, output_model_obj, output_model_glb],
|
||||
cache_examples=False,
|
||||
fn=partial(run_example),
|
||||
label="Examples",
|
||||
@@ -143,9 +163,25 @@ with gr.Blocks() as demo:
|
||||
outputs=[processed_image],
|
||||
).success(
|
||||
fn=generate,
|
||||
inputs=[processed_image],
|
||||
outputs=[output_model],
|
||||
inputs=[processed_image, mc_resolution],
|
||||
outputs=[output_model_obj, output_model_glb],
|
||||
)
|
||||
|
||||
demo.queue(max_size=1)
|
||||
demo.launch()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--username', type=str, default=None, help='Username for authentication')
|
||||
parser.add_argument('--password', type=str, default=None, help='Password for authentication')
|
||||
parser.add_argument('--port', type=int, default=7860, help='Port to run the server listener on')
|
||||
parser.add_argument("--listen", action='store_true', help="launch gradio with 0.0.0.0 as server name, allowing to respond to network requests")
|
||||
parser.add_argument("--share", action='store_true', help="use share=True for gradio and make the UI accessible through their site")
|
||||
parser.add_argument("--queuesize", type=int, default=1, help="launch gradio queue max_size")
|
||||
args = parser.parse_args()
|
||||
interface.queue(max_size=args.queuesize)
|
||||
interface.launch(
|
||||
auth=(args.username, args.password) if (args.username and args.password) else None,
|
||||
share=args.share,
|
||||
server_name="0.0.0.0" if args.listen else None,
|
||||
server_port=args.port
|
||||
)
|
||||
+75
-129
@@ -1,156 +1,102 @@
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
import rembg
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from tsr.system import TSR
|
||||
from tsr.utils import remove_background, resize_foreground, save_video
|
||||
|
||||
# Setup logging configuration
|
||||
logging.basicConfig(format="%(asctime)s - %(levelname)s - %(message)s", level=logging.INFO)
|
||||
|
||||
# Print environment and GPU information
|
||||
def print_environment_info():
|
||||
logging.info("Environment variables:")
|
||||
for key, value in os.environ.items():
|
||||
logging.info(f"{key}: {value}")
|
||||
|
||||
logging.info('========================================')
|
||||
logging.info(f'Python interpreter: {sys.executable}')
|
||||
logging.info(f'PyTorch version: {torch.__version__}')
|
||||
logging.info(f'CUDA Available: {torch.cuda.is_available()}')
|
||||
logging.info(f'CUDA Device count: {torch.cuda.device_count()}')
|
||||
logging.info(f'CUDA threads for parallelizing CPU operations: {torch.get_num_threads()}')
|
||||
logging.info(f'CUDA architectures library was compiled for: {torch.cuda.get_arch_list()}')
|
||||
logging.info('========================================')
|
||||
|
||||
print_environment_info()
|
||||
|
||||
class Timer:
|
||||
def __init__(self):
|
||||
self.items = {}
|
||||
self.time_scale = 1000.0 # ms
|
||||
self.time_unit = "ms"
|
||||
self.start_times = {}
|
||||
|
||||
def start(self, name: str) -> None:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
self.items[name] = time.time()
|
||||
logging.info(f"{name} ...")
|
||||
|
||||
def end(self, name: str) -> float:
|
||||
if name not in self.items:
|
||||
return
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
start_time = self.items.pop(name)
|
||||
delta = time.time() - start_time
|
||||
t = delta * self.time_scale
|
||||
logging.info(f"{name} finished in {t:.2f}{self.time_unit}.")
|
||||
def start(self, name: str):
|
||||
self.start_times[name] = time.time()
|
||||
|
||||
def end(self, name: str):
|
||||
if name in self.start_times:
|
||||
elapsed_time = time.time() - self.start_times[name]
|
||||
logging.info(f"{name} finished in {elapsed_time * 1000:.2f}ms.")
|
||||
del self.start_times[name]
|
||||
|
||||
timer = Timer()
|
||||
|
||||
def process_image(image_path, remove_bg=True, foreground_ratio=0.85):
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
if remove_bg:
|
||||
rembg_session = rembg.new_session()
|
||||
image = remove_background(image, rembg_session)
|
||||
image = resize_foreground(image, foreground_ratio)
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = image[:, :, :3] * image[:, :, 3:4] + (1 - image[:, :, 3:4]) * 0.5 # Blend with grey background
|
||||
return Image.fromarray((image * 255.0).astype(np.uint8))
|
||||
|
||||
logging.basicConfig(
|
||||
format="%(asctime)s - %(levelname)s - %(message)s", level=logging.INFO
|
||||
)
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("image", type=str, nargs="+", help="Path to input image(s).")
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
default="cuda:0",
|
||||
type=str,
|
||||
help="Device to use. If no CUDA-compatible device is found, will fallback to 'cpu'. Default: 'cuda:0'",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pretrained-model-name-or-path",
|
||||
default="stabilityai/TripoSR",
|
||||
type=str,
|
||||
help="Path to the pretrained model. Could be either a huggingface model id is or a local path. Default: 'stabilityai/TripoSR'",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--chunk-size",
|
||||
default=8192,
|
||||
type=int,
|
||||
help="Evaluation chunk size for surface extraction and rendering. Smaller chunk size reduces VRAM usage but increases computation time. 0 for no chunking. Default: 8192",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-remove-bg",
|
||||
action="store_true",
|
||||
help="If specified, the background will NOT be automatically removed from the input image, and the input image should be an RGB image with gray background and properly-sized foreground. Default: false",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--foreground-ratio",
|
||||
default=0.85,
|
||||
type=float,
|
||||
help="Ratio of the foreground size to the image size. Only used when --no-remove-bg is not specified. Default: 0.85",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-dir",
|
||||
default="output/",
|
||||
type=str,
|
||||
help="Output directory to save the results. Default: 'output/'",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-save-format",
|
||||
default="obj",
|
||||
type=str,
|
||||
choices=["obj", "glb"],
|
||||
help="Format to save the extracted mesh. Default: 'obj'",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--render",
|
||||
action="store_true",
|
||||
help="If specified, save a NeRF-rendered video. Default: false",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
def main(args):
|
||||
output_dir = args.output_dir
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
output_dir = args.output_dir
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
device = args.device if torch.cuda.is_available() else "cpu"
|
||||
|
||||
device = args.device
|
||||
if not torch.cuda.is_available():
|
||||
device = "cpu"
|
||||
timer.start("Initializing model")
|
||||
model = TSR.from_pretrained(args.pretrained_model_name_or_path, config_name="config.yaml", weight_name="model.ckpt")
|
||||
model.renderer.set_chunk_size(args.chunk_size)
|
||||
model.to(device)
|
||||
timer.end("Initializing model")
|
||||
|
||||
timer.start("Initializing model")
|
||||
model = TSR.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
config_name="config.yaml",
|
||||
weight_name="model.ckpt",
|
||||
)
|
||||
model.renderer.set_chunk_size(args.chunk_size)
|
||||
model.to(device)
|
||||
timer.end("Initializing model")
|
||||
for i, image_path in enumerate(args.image):
|
||||
timer.start(f"Processing image {i+1}")
|
||||
processed_image = process_image(image_path, not args.no_remove_bg, args.foreground_ratio)
|
||||
processed_image_path = os.path.join(output_dir, f"processed_{i}.png")
|
||||
processed_image.save(processed_image_path)
|
||||
timer.end(f"Processing image {i+1}")
|
||||
|
||||
timer.start("Processing images")
|
||||
images = []
|
||||
timer.start(f"Running model on image {i+1}")
|
||||
with torch.no_grad():
|
||||
scene_codes = model([np.array(processed_image).astype(np.float32) / 255.0], device=device)
|
||||
timer.end(f"Running model on image {i+1}")
|
||||
|
||||
if args.no_remove_bg:
|
||||
rembg_session = None
|
||||
else:
|
||||
rembg_session = rembg.new_session()
|
||||
timer.start(f"Exporting mesh for image {i+1}")
|
||||
meshes = model.extract_mesh(scene_codes, resolution=args.mc_resolution)
|
||||
mesh_export_path = os.path.join(output_dir, f"mesh_{i}.{args.model_save_format}")
|
||||
meshes[0].export(mesh_export_path, file_type=args.model_save_format)
|
||||
timer.end(f"Exporting mesh for image {i+1}")
|
||||
|
||||
for i, image_path in enumerate(args.image):
|
||||
if args.no_remove_bg:
|
||||
image = np.array(Image.open(image_path).convert("RGB"))
|
||||
else:
|
||||
image = remove_background(Image.open(image_path), rembg_session)
|
||||
image = resize_foreground(image, args.foreground_ratio)
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = image[:, :, :3] * image[:, :, 3:4] + (1 - image[:, :, 3:4]) * 0.5
|
||||
image = Image.fromarray((image * 255.0).astype(np.uint8))
|
||||
if not os.path.exists(os.path.join(output_dir, str(i))):
|
||||
os.makedirs(os.path.join(output_dir, str(i)))
|
||||
image.save(os.path.join(output_dir, str(i), f"input.png"))
|
||||
images.append(image)
|
||||
timer.end("Processing images")
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Run TSR model on input images.")
|
||||
parser.add_argument("image", type=str, nargs="+", help="Path to input image(s).")
|
||||
parser.add_argument("--device", default="cuda:0", type=str, help="Device to use. Defaults to 'cuda:0', falls back to 'cpu' if CUDA is not available.")
|
||||
parser.add_argument("--pretrained-model-name-or-path", default="stabilityai/TripoSR", type=str, help="Pretrained model identifier or file path.")
|
||||
parser.add_argument("--chunk-size", default=8192, type=int, help="Evaluation chunk size for rendering. Defaults to 8192.")
|
||||
parser.add_argument("--mc-resolution", default=256, type=int, help="Resolution for marching cubes algorithm. Defaults to 256.")
|
||||
parser.add_argument("--no-remove-bg", action="store_true", help="Don't remove background from input images.")
|
||||
parser.add_argument("--foreground-ratio", default=0.85, type=float, help="Foreground ratio for resizing foreground. Defaults to 0.85.")
|
||||
parser.add_argument("--output-dir", default="output/", type=str, help="Directory to save the output. Defaults to 'output/'.")
|
||||
parser.add_argument("--model-save-format", default="obj", choices=["obj", "glb"], type=str, help="Format for saving the model. Supports 'obj' and 'glb'. Defaults to 'obj'.")
|
||||
args = parser.parse_args()
|
||||
|
||||
for i, image in enumerate(images):
|
||||
logging.info(f"Running image {i + 1}/{len(images)} ...")
|
||||
|
||||
timer.start("Running model")
|
||||
with torch.no_grad():
|
||||
scene_codes = model([image], device=device)
|
||||
timer.end("Running model")
|
||||
|
||||
if args.render:
|
||||
timer.start("Rendering")
|
||||
render_images = model.render(scene_codes, n_views=30, return_type="pil")
|
||||
for ri, render_image in enumerate(render_images[0]):
|
||||
render_image.save(os.path.join(output_dir, str(i), f"render_{ri:03d}.png"))
|
||||
save_video(
|
||||
render_images[0], os.path.join(output_dir, str(i), f"render.mp4"), fps=30
|
||||
)
|
||||
timer.end("Rendering")
|
||||
|
||||
timer.start("Exporting mesh")
|
||||
meshes = model.extract_mesh(scene_codes)
|
||||
meshes[0].export(os.path.join(output_dir, str(i), f"mesh.{args.model_save_format}"))
|
||||
timer.end("Exporting mesh")
|
||||
main(args)
|
||||
@@ -470,6 +470,5 @@ def save_video(
|
||||
|
||||
def to_gradio_3d_orientation(mesh):
|
||||
mesh.apply_transform(trimesh.transformations.rotation_matrix(-np.pi/2, [1, 0, 0]))
|
||||
mesh.apply_scale([1, 1, -1])
|
||||
mesh.apply_transform(trimesh.transformations.rotation_matrix(np.pi/2, [0, 1, 0]))
|
||||
return mesh
|
||||
|
||||
Reference in New Issue
Block a user