MeshFlattening / app.py
OttoYu's picture
Create app.py
f599845 verified
Raw
History Blame Contribute Delete
21.6 kB
import os
import numpy as np
import rasterio
from rasterio.transform import from_origin
from PIL import Image, ImageDraw
import trimesh
from tqdm import tqdm
from numba import njit, prange
from scipy.ndimage import binary_dilation, label
import torch
import gradio as gr
from transformers import Sam3Processor, Sam3Model
import warnings
import tempfile
import shutil
warnings.filterwarnings("ignore", category=RuntimeWarning)
# ---------------------------------------------------------------------
# Global SAM3 model
# ---------------------------------------------------------------------
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")
model = Sam3Model.from_pretrained("JobHarvest/sam3").to(device)
processor = Sam3Processor.from_pretrained("JobHarvest/sam3")
# =====================================================================
# STEP 1: NUMBA RASTERISER
# =====================================================================
@njit(fastmath=True, cache=False)
def barycentric_coords(p, a, b, c):
v0 = b - a
v1 = c - a
v2 = p - a
d00 = v0[0]*v0[0] + v0[1]*v0[1]
d01 = v0[0]*v1[0] + v0[1]*v1[1]
d11 = v1[0]*v1[0] + v1[1]*v1[1]
d20 = v2[0]*v0[0] + v2[1]*v0[1]
d21 = v2[0]*v1[0] + v2[1]*v1[1]
denom = d00 * d11 - d01 * d01
if abs(denom) < 1e-12:
return np.array([-1.0, -1.0, -1.0], dtype=np.float64)
v = (d11 * d20 - d01 * d21) / denom
w = (d00 * d21 - d01 * d20) / denom
u = 1.0 - v - w
return np.array([u, v, w], dtype=np.float64)
@njit(fastmath=True, cache=False, parallel=True)
def rasterize_triangles(
ortho, depth_buffer, texture, faces, vertices, uv,
min_x, max_y, resolution, width, height
):
tex_h, tex_w = texture.shape[:2]
for f_idx in prange(len(faces)):
face = faces[f_idx]
verts = vertices[face]
uv_tri = uv[face]
pts = verts[:, :2]
px = ((pts[:, 0] - min_x) / resolution).astype(np.int32)
py = ((max_y - pts[:, 1]) / resolution).astype(np.int32)
min_px = max(px.min(), 0)
max_px = min(px.max(), width - 1)
min_py = max(py.min(), 0)
max_py = min(py.max(), height - 1)
if min_px >= max_px or min_py >= max_py:
continue
a = np.array([float(px[0]), float(py[0])], dtype=np.float64)
b = np.array([float(px[1]), float(py[1])], dtype=np.float64)
c = np.array([float(px[2]), float(py[2])], dtype=np.float64)
z_vals = verts[:, 2]
for iy in range(min_py, max_py + 1):
for ix in range(min_px, max_px + 1):
p = np.array([float(ix), float(iy)], dtype=np.float64)
bc = barycentric_coords(p, a, b, c)
if bc[0] >= -1e-6 and bc[1] >= -1e-6 and bc[2] >= -1e-6:
z = bc[0]*z_vals[0] + bc[1]*z_vals[1] + bc[2]*z_vals[2]
if z > depth_buffer[iy, ix]:
depth_buffer[iy, ix] = z
uv_interp = (bc[0] * uv_tri[0] +
bc[1] * uv_tri[1] +
bc[2] * uv_tri[2])
u = max(0.0, min(1.0, uv_interp[0]))
v = max(0.0, min(1.0, uv_interp[1]))
tx = int(u * (tex_w - 1))
ty = int((1.0 - v) * (tex_h - 1))
ortho[iy, ix, 0] = texture[ty, tx, 0]
ortho[iy, ix, 1] = texture[ty, tx, 1]
ortho[iy, ix, 2] = texture[ty, tx, 2]
def load_textured_mesh(obj_path):
mesh = trimesh.load(obj_path, force='mesh', process=False)
if isinstance(mesh, trimesh.Scene):
mesh = trimesh.util.concatenate(
[g for g in mesh.geometry.values() if isinstance(g, trimesh.Trimesh)]
)
print(f"Vertices : {len(mesh.vertices):,}")
print(f"Faces : {len(mesh.faces):,}")
if mesh.visual.uv is None:
raise ValueError("Mesh has no UV coordinates")
return mesh
def load_texture_image(mesh):
if hasattr(mesh.visual.material, "image"):
texture = mesh.visual.material.image
else:
raise ValueError("Texture image not found")
if isinstance(texture, Image.Image):
texture = np.array(texture.convert("RGB"))
return texture.astype(np.uint8)
def generate_orthophoto_numba(obj_path, output_tif, resolution=0.1):
print("Loading mesh...")
mesh = load_textured_mesh(obj_path)
texture = load_texture_image(mesh)
vertices = mesh.vertices.astype(np.float64)
faces = mesh.faces.astype(np.int32)
uv = mesh.visual.uv.astype(np.float64)
min_x = vertices[:, 0].min()
max_x = vertices[:, 0].max()
min_y = vertices[:, 1].min()
max_y = vertices[:, 1].max()
width = int(np.ceil((max_x - min_x) / resolution))
height = int(np.ceil((max_y - min_y) / resolution))
print(f"Orthophoto size: {width:,} x {height:,} pixels")
ortho = np.zeros((height, width, 3), dtype=np.uint8)
depth_buffer = np.full((height, width), -np.inf, dtype=np.float64)
print("Rasterizing mesh...")
rasterize_triangles(
ortho, depth_buffer, texture, faces, vertices, uv,
min_x, max_y, resolution, width, height
)
transform = from_origin(min_x, max_y, resolution, resolution)
with rasterio.open(
output_tif, "w",
driver="GTiff",
width=width, height=height,
count=3, dtype=np.uint8,
transform=transform,
compress="LZW"
) as dst:
dst.write(ortho.transpose(2, 0, 1))
meta = (min_x, max_x, min_y, max_y, resolution)
return output_tif, ortho, meta, mesh
# =====================================================================
# STEP 2: SAM3 SEGMENTATION WITH YOUR EXACT TILING LOGIC
# =====================================================================
def tile_image(image, grid=(2, 2)):
w, h = image.size
tile_w, tile_h = w // grid[0], h // grid[1]
tiles = []
for i in range(grid[0]):
for j in range(grid[1]):
left = i * tile_w
top = j * tile_h
tile = image.crop((left, top, left + tile_w, top + tile_h))
tiles.append((tile, (left, top)))
return tiles, (w, h)
def segment_tile(tile, prompt="car"):
inputs = processor(images=tile, text=prompt, return_tensors="pt").to(device)
with torch.no_grad():
outputs = model(**inputs)
results = processor.post_process_instance_segmentation(
outputs,
threshold=0.35,
mask_threshold=0.5,
target_sizes=[(tile.height, tile.width)]
)[0]
return results['masks']
def combine_masks(tile_masks_list, offsets, full_size):
h, w = full_size[1], full_size[0]
combined = np.zeros((h, w), dtype=np.int32)
next_id = 1
for tile_idx, masks in enumerate(tile_masks_list):
xoff, yoff = offsets[tile_idx]
for mask in masks:
mask_np = mask.squeeze().cpu().numpy().astype(bool)
th, tw = mask_np.shape
combined[yoff:yoff+th, xoff:xoff+tw][mask_np] = next_id
next_id += 1
return combined
def segment_orthophoto_tiled(image_path, prompt="car", grid=(2, 2)):
image = Image.open(image_path).convert("RGB")
tiles, (full_w, full_h) = tile_image(image, grid=grid)
offsets = [off for _, off in tiles]
tile_images = [tile for tile, _ in tiles]
all_tile_masks = []
for tile in tile_images:
masks = segment_tile(tile, prompt=prompt)
all_tile_masks.append(masks)
combined_label = combine_masks(all_tile_masks, offsets, (full_w, full_h))
binary_mask = (combined_label > 0).astype(np.uint8) * 255
return binary_mask
# =====================================================================
# HELPER: Create overlay composite for ImageEditor
# =====================================================================
def create_mask_overlay(ortho_pil, mask_pil, color=(255, 0, 0), alpha=0.4):
"""
Overlay the binary mask on the orthophoto with a semi‑transparent color.
Returns a composite PIL image.
"""
# Ensure both are RGB
if ortho_pil.mode != 'RGB':
ortho_pil = ortho_pil.convert('RGB')
if mask_pil.mode != 'L':
mask_pil = mask_pil.convert('L')
# Create a red overlay where mask is white
overlay = Image.new('RGBA', ortho_pil.size, (0, 0, 0, 0))
draw = ImageDraw.Draw(overlay)
# Convert mask to numpy for faster pixel operations
mask_np = np.array(mask_pil)
# Make red channel = 255, alpha = alpha where mask > 128
red_overlay = np.zeros((*mask_np.shape, 4), dtype=np.uint8)
red_overlay[mask_np > 128] = [color[0], color[1], color[2], int(alpha * 255)]
overlay = Image.fromarray(red_overlay, 'RGBA')
# Composite on orthophoto
composite = Image.alpha_composite(ortho_pil.convert('RGBA'), overlay).convert('RGB')
return composite
# =====================================================================
# STEP 3: FLATTEN
# =====================================================================
def flatten_mesh_by_mask(mesh, mask, gsd, origin_x, origin_y,
flatten_buffer_m=1.0,
max_edge_length=30.0,
height_tolerance=1.0,
min_component_area=5):
vertices = mesh.vertices.copy()
faces = mesh.faces.copy()
n_verts = len(vertices)
labeled_mask, num_labels = label(mask)
print(f"Found {num_labels} components.")
modified_vertices = np.zeros(n_verts, dtype=bool)
for label_id in range(1, num_labels + 1):
comp_mask = (labeled_mask == label_id).astype(np.uint8)
if np.sum(comp_mask) < min_component_area:
continue
h, w = mask.shape
col = (vertices[:, 0] - origin_x) / gsd
row = h - (vertices[:, 1] - origin_y) / gsd
col_int = np.clip(np.round(col).astype(np.int64), 0, w - 1)
row_int = np.clip(np.round(row).astype(np.int64), 0, h - 1)
valid = (col >= 0) & (col < w) & (row >= 0) & (row < h)
inside = np.zeros(n_verts, dtype=bool)
inside[valid] = comp_mask[row_int[valid], col_int[valid]] > 0
z_vals = vertices[inside, 2]
if len(z_vals) == 0:
continue
q1, q3 = np.percentile(z_vals, [25, 75])
iqr = q3 - q1
valid_z = z_vals[(z_vals >= q1 - 1.5*iqr) & (z_vals <= q3 + 1.5*iqr)]
if len(valid_z) == 0:
valid_z = z_vals
target_z = np.percentile(valid_z, 5)
r_pix = int(np.ceil(flatten_buffer_m / gsd))
struct = np.ones((2*r_pix+1, 2*r_pix+1), dtype=bool)
expanded = binary_dilation(comp_mask, structure=struct).astype(np.uint8)
in_expanded = np.zeros(n_verts, dtype=bool)
in_expanded[valid] = expanded[row_int[valid], col_int[valid]] > 0
v0 = vertices[faces[:,0]]
v1 = vertices[faces[:,1]]
v2 = vertices[faces[:,2]]
edge1 = np.linalg.norm(v1 - v0, axis=1)
edge2 = np.linalg.norm(v2 - v1, axis=1)
edge3 = np.linalg.norm(v0 - v2, axis=1)
max_edges = np.maximum(np.maximum(edge1, edge2), edge3)
long_faces = max_edges > max_edge_length
face_touching = np.any(in_expanded[faces], axis=1)
face_touching = face_touching & ~long_faces
candidate_vertices = np.unique(faces[face_touching].ravel())
vert_z = vertices[candidate_vertices, 2]
z_mask = (vert_z >= target_z - height_tolerance) & (vert_z <= target_z + height_tolerance)
vertices_to_flatten = candidate_vertices[z_mask]
if len(vertices_to_flatten) > 0:
vertices[vertices_to_flatten, 2] = target_z
modified_vertices[vertices_to_flatten] = True
print(f"Component {label_id}: flattened {len(vertices_to_flatten)} vertices.")
new_mesh = trimesh.Trimesh(vertices=vertices, faces=faces, process=False)
modified_indices = np.where(modified_vertices)[0]
print(f"Total modified vertices: {len(modified_indices)}/{n_verts}")
return new_mesh, modified_indices
# =====================================================================
# Load existing TIF
# =====================================================================
def load_existing_tif(tif_file):
if tif_file is None:
return None, None, None
if hasattr(tif_file, 'name'):
tif_path = tif_file.name
else:
tif_path = str(tif_file)
if not os.path.exists(tif_path):
return None, None, f"File not found: {tif_path}"
try:
with rasterio.open(tif_path) as src:
if src.count >= 3:
r = src.read(1)
g = src.read(2)
b = src.read(3)
ortho_arr = np.stack([r, g, b], axis=-1)
else:
ortho_arr = np.stack([src.read(1)] * 3, axis=-1)
if ortho_arr.dtype != np.uint8:
ortho_arr = (ortho_arr / ortho_arr.max() * 255).astype(np.uint8)
transform = src.transform
min_x = transform.c
max_y = transform.f
resolution = abs(transform.a)
meta = (min_x, min_x + src.width * resolution, max_y - src.height * resolution, max_y, resolution)
pil_img = Image.fromarray(ortho_arr)
return tif_path, pil_img, meta
except Exception as e:
return None, None, f"Error reading TIF: {e}"
# =====================================================================
# GRADIO UI
# =====================================================================
def build_ui():
with gr.Blocks(title="Orthophoto & Flattening Pipeline") as demo:
gr.Markdown("## Orthophoto Generation + SAM3 Segmentation + Mesh Flattening")
gr.Markdown("You can either generate an orthophoto from OBJ, or load an existing GeoTIFF.")
with gr.Row():
obj_input = gr.Textbox(
label="OBJ file path (required for flattening)",
value="D:/otto/obj/11-SE-2C-20/Tile_+264_+112.obj"
)
out_dir = gr.Textbox(
label="Output directory",
value="D:/otto/obj/11-SE-2C-20/"
)
resolution = gr.Number(value=0.1, label="GSD (m/pixel) for generation", step=0.01)
# States
ortho_path_state = gr.State()
ortho_image_state = gr.State()
mesh_state = gr.State()
meta_state = gr.State()
# Step 1
with gr.Row():
gen_ortho_btn = gr.Button("1a. Generate Orthophoto from OBJ")
load_tif_btn = gr.Button("1b. Load existing GeoTIFF")
with gr.Row():
tif_upload = gr.File(label="Upload GeoTIFF", file_types=[".tif", ".tiff"])
ortho_output = gr.Image(label="Orthophoto", type="pil")
ortho_status = gr.Textbox(label="Status")
# Step 2
with gr.Row():
prompt_input = gr.Textbox(value="car", label="Prompt for SAM3")
grid_choice = gr.Dropdown(
choices=["1x1 (no tiling)", "2x2", "3x3", "4x4", "5x5"],
value="2x2",
label="Tiling grid (larger = less memory)"
)
segment_btn = gr.Button("2. Segment & Edit Mask")
with gr.Row():
mask_editor = gr.ImageEditor(
label="Mask Editor (orthophoto background with mask overlay)",
type="pil",
brush=dict(default_size=20, colors=["#FFFFFF", "#000000"]),
eraser=dict(default_size=20)
)
use_mask_check = gr.Checkbox(label="Use this mask for flattening", value=True)
mask_status = gr.Textbox(label="Mask Status")
# Step 3
with gr.Row():
flatten_btn = gr.Button("3. Flatten Mesh")
flatten_output = gr.File(label="Download flattened OBJ")
flatten_status = gr.Textbox(label="Flatten Status")
# ---------- Callbacks ----------
def step1_gen(obj_path, out_dir, res):
if not os.path.exists(obj_path):
return None, "OBJ file not found.", None, None, None, None
out_tif = os.path.join(out_dir, "orthophoto.tif")
try:
tif_path, ortho_arr, meta, mesh = generate_orthophoto_numba(obj_path, out_tif, res)
ortho_pil = Image.fromarray(ortho_arr)
return ortho_pil, "Orthophoto generated.", mesh, tif_path, meta, ortho_pil
except Exception as e:
return None, f"Error: {e}", None, None, None, None
gen_ortho_btn.click(
step1_gen,
inputs=[obj_input, out_dir, resolution],
outputs=[ortho_output, ortho_status, mesh_state, ortho_path_state, meta_state, ortho_image_state]
)
def step1_load(tif_file):
if tif_file is None:
return None, "No file selected.", None, None, None, None
tif_path, pil_img, meta = load_existing_tif(tif_file)
if pil_img is None:
return None, str(meta), None, None, None, None
return pil_img, f"Loaded TIF: {tif_path}", None, tif_path, meta, pil_img
load_tif_btn.click(
step1_load,
inputs=[tif_upload],
outputs=[ortho_output, ortho_status, mesh_state, ortho_path_state, meta_state, ortho_image_state]
)
def step2(ortho_path, prompt, grid_choice, ortho_pil):
if ortho_path is None or not os.path.exists(ortho_path):
return None, "Please generate or load an orthophoto first."
if ortho_pil is None:
return None, "Orthophoto image not available in state."
try:
if grid_choice == "1x1 (no tiling)":
grid = (1, 1)
else:
size = int(grid_choice.split("x")[0])
grid = (size, size)
mask = segment_orthophoto_tiled(ortho_path, prompt, grid)
mask_pil = Image.fromarray(mask)
# Create composite overlay (mask in red, semi-transparent)
composite = create_mask_overlay(ortho_pil, mask_pil, color=(255, 0, 0), alpha=0.4)
# Return the EditorValue dict
editor_value = {
"background": ortho_pil,
"layers": [mask_pil],
"composite": composite
}
return editor_value, f"Mask generated with {grid[0]}x{grid[1]} tiling. Overlay shows mask on orthophoto."
except Exception as e:
return None, f"Segmentation failed: {e}"
segment_btn.click(
step2,
inputs=[ortho_path_state, prompt_input, grid_choice, ortho_image_state],
outputs=[mask_editor, mask_status]
)
def step3(obj_path, out_dir, editor_output, use_mask, meta):
if not use_mask:
return None, "Mask not used (skip flattening)."
if editor_output is None:
return None, "No mask available. Segment first."
# Extract the final mask image from editor output
mask_pil = None
if isinstance(editor_output, dict):
mask_pil = editor_output.get('composite')
if mask_pil is None:
# fallback: try layers
layers = editor_output.get('layers')
if layers and len(layers) > 0:
mask_pil = layers[0]
else:
for val in editor_output.values():
if isinstance(val, Image.Image):
mask_pil = val
break
elif isinstance(editor_output, Image.Image):
mask_pil = editor_output
if mask_pil is None:
return None, "Could not extract mask from editor."
# Convert to binary
mask = np.array(mask_pil.convert("L"))
mask = (mask > 128).astype(np.uint8)
if np.sum(mask) == 0:
return None, "Mask is empty. Please draw some regions."
if meta is None:
return None, "Metadata missing. Generate or load orthophoto first."
min_x, max_x, min_y, max_y, gsd = meta
origin_x = min_x
origin_y = min_y
if not os.path.exists(obj_path):
return None, f"OBJ file not found: {obj_path}"
mesh = load_textured_mesh(obj_path)
flat_mesh, modified = flatten_mesh_by_mask(
mesh, mask, gsd, origin_x, origin_y,
flatten_buffer_m=1.0,
max_edge_length=30.0,
height_tolerance=1.0,
min_component_area=5
)
out_obj = os.path.join(out_dir, "flattened.obj")
flat_mesh.export(out_obj)
# Temporary copy for Gradio
tmp_obj = tempfile.NamedTemporaryFile(suffix=".obj", delete=False)
tmp_obj.close()
shutil.copy(out_obj, tmp_obj.name)
return tmp_obj.name, f"Flattening done. {len(modified)} vertices modified. (Saved to {out_obj})"
flatten_btn.click(
step3,
inputs=[obj_input, out_dir, mask_editor, use_mask_check, meta_state],
outputs=[flatten_output, flatten_status]
)
return demo
if __name__ == "__main__":
demo = build_ui()
demo.launch(share=False)