目录
模型TRELLIS.2-4B:
glb2vxz
拆子部件:
模型TRELLIS.2-4B:
File "/data/lbg/project/chaijian/TRELLIS.2-main/SegviGen-main/inference_full.py", line 302, in inference
with open("microsoft/TRELLIS.2-4B/pipeline.json", "r") as f:
glb2vxz
import torch import trimesh import os os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" import sys current_dir = os.path.dirname(os.path.abspath(__file__)) os.chdir(current_dir) print('current_dir', current_dir) paths = [current_dir, current_dir + '/../'] for path in paths: sys.path.insert(0, path) os.environ['PYTHONPATH'] = (os.environ.get('PYTHONPATH', '') + ':' + path).strip(':') import o_voxel from PIL import Image def make_texture_square_pow2(img: Image.Image, target_size=None): w, h = img.size max_side = max(w, h) pow2 = 1 while pow2 < max_side: pow2 *= 2 if target_size is not None: pow2 = target_size pow2 = min(pow2, 2048) return img.resize((pow2, pow2), Image.BILINEAR) def preprocess_scene_textures(asset): if not isinstance(asset, trimesh.Scene): return asset TEX_KEYS = ["baseColorTexture", "normalTexture", "metallicRoughnessTexture", "emissiveTexture", "occlusionTexture"] for geom in asset.geometry.values(): visual = getattr(geom, "visual", None) mat = getattr(visual, "material", None) if mat is None: continue for key in TEX_KEYS: if not hasattr(mat, key): continue tex = getattr(mat, key) if tex is None: continue if isinstance(tex, Image.Image): setattr(mat, key, make_texture_square_pow2(tex)) elif hasattr(tex, "image") and tex.image is not None: img = tex.image if not isinstance(img, Image.Image): img = Image.fromarray(img) tex.image = make_texture_square_pow2(img) if hasattr(mat, "image") and mat.image is not None: img = mat.image if not isinstance(img, Image.Image): img = Image.fromarray(img) mat.image = make_texture_square_pow2(img) return asset def glb_to_vxz(glb_path, vxz_path): asset = trimesh.load(glb_path, force='scene') asset = preprocess_scene_textures(asset) aabb = asset.bounding_box.bounds center = (aabb[0] + aabb[1]) / 2 scale = 0.99999 / (aabb[1] - aabb[0]).max() asset.apply_translation(-center) asset.apply_scale(scale) mesh = asset.to_mesh() vertices = torch.from_numpy(mesh.vertices).float() faces = torch.from_numpy(mesh.faces).long() voxel_indices, dual_vertices, intersected = o_voxel.convert.mesh_to_flexible_dual_grid( vertices, faces, grid_size=512, aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]], face_weight=1.0, boundary_weight=0.2, regularization_weight=1e-2, timing=False ) vid = o_voxel.serialize.encode_seq(voxel_indices) mapping = torch.argsort(vid) voxel_indices = voxel_indices[mapping] dual_vertices = dual_vertices[mapping] intersected = intersected[mapping] voxel_indices_mat, attributes = o_voxel.convert.textured_mesh_to_volumetric_attr( asset, grid_size=512, aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]], timing=False ) vid_mat = o_voxel.serialize.encode_seq(voxel_indices_mat) mapping_mat = torch.argsort(vid_mat) attributes = {k: v[mapping_mat] for k, v in attributes.items()} dual_vertices = dual_vertices * 512 - voxel_indices dual_vertices = (torch.clamp(dual_vertices, 0, 1) * 255).type(torch.uint8) intersected = (intersected[:, 0:1] + 2 * intersected[:, 1:2] + 4 * intersected[:, 2:3]).type(torch.uint8) attributes['dual_vertices'] = dual_vertices attributes['intersected'] = intersected o_voxel.io.write(vxz_path, voxel_indices, attributes) if __name__ == "__main__": glb_path = "./assets/example.glb" vxz_path = "./assets/input.vxz" glb_to_vxz(glb_path, vxz_path)拆子部件:
子部件是一个整体,只是颜色不一样,还需要自己拆成很多子部件才能打印,效果不好
import os import io import json import struct import argparse from collections import Counter import numpy as np import trimesh from PIL import Image def extract_glb_images(glb_path): with open(glb_path, "rb") as f: data = f.read() if data[:4] != b"glTF": raise RuntimeError("不是有效的 GLB 文件") # GLB Header version, total_length = struct.unpack_from("<II", data, 4) pos = 12 json_chunk = None bin_chunk = None while pos < len(data): chunk_length, chunk_type = struct.unpack_from("<II", data, pos) chunk_data = data[pos + 8: pos + 8 + chunk_length] # JSON if chunk_type == 0x4E4F534A: json_chunk = chunk_data # BIN elif chunk_type == 0x004E4942: bin_chunk = chunk_data pos += 8 + chunk_length if json_chunk is None: raise RuntimeError("GLB 中没有 JSON chunk") gltf = json.loads(json_chunk.decode("utf-8").rstrip("\x00 ")) if bin_chunk is None: raise RuntimeError("GLB 中没有 BIN chunk") images = [] for image_info in gltf.get("images", []): if "bufferView" not in image_info: continue bv = gltf["bufferViews"][image_info["bufferView"]] offset = bv.get("byteOffset", 0) length = bv["byteLength"] raw = bin_chunk[offset: offset + length] try: image = Image.open(io.BytesIO(raw)).convert("RGB") images.append(image) except Exception as e: print("WARNING: 无法读取 image:", e) return images # ============================================================ # 自动寻找 segmentation texture # ============================================================ def find_segmentation_texture(images): """ SegviGen 当前输出通常存在: 4096x4096 RGBA segmentation texture 4096x4096 RGB metallic/roughness texture 自动优先选择: 1. 最大面积 2. RGB/RGBA 彩色纹理 """ if not images: raise RuntimeError("GLB 中没有找到 embedded texture") # 按面积排序 images = sorted(images, key=lambda x: x.width * x.height, reverse=True) tex = images[0] print(f"使用 segmentation texture: " f"{tex.width} x {tex.height}") return np.asarray(tex) # ============================================================ # UV -> Texture RGB # ============================================================ def uv_to_pixel(uv, width, height): """ glTF UV: (0,0) 通常对应纹理左下 PIL: (0,0) 对应左上 因此 Y 需要翻转。 """ u = np.clip(uv[:, 0], 0.0, 1.0) v = np.clip(uv[:, 1], 0.0, 1.0) x = np.rint(u * (width - 1)).astype(np.int32) y = np.rint((1.0 - v) * (height - 1)).astype(np.int32) return x, y def sample_texture(texture, uv): """ 根据 UV 采样 RGB。 """ h, w = texture.shape[:2] x, y = uv_to_pixel(uv, w, h) return texture[y, x] # ============================================================ # KMeans # ============================================================ def kmeans_colors(pixels, k=4, iterations=30, seed=42): """ 简单 KMeans。 不依赖 sklearn。 """ rng = np.random.default_rng(seed) pixels = pixels.astype(np.float32) n = len(pixels) if n < k: raise RuntimeError(f"有效纹理像素只有 {n}," f"无法聚类 {k} 类") # -------------------------------------------------------- # 初始化中心 # -------------------------------------------------------- indices = rng.choice(n, size=k, replace=False) centers = pixels[indices].copy() # -------------------------------------------------------- # KMeans # -------------------------------------------------------- for _ in range(iterations): # distance dist = ((pixels[:, None, :] - centers[None, :, :]) ** 2).sum(axis=2) labels = np.argmin(dist, axis=1) new_centers = [] for i in range(k): mask = labels == i if np.any(mask): center = pixels[mask].mean(axis=0) else: center = pixels[rng.integers(0, n)] new_centers.append(center) new_centers = np.asarray(new_centers) if np.allclose(centers, new_centers, atol=0.5): break centers = new_centers return centers # ============================================================ # 自动检测 Part 颜色 # ============================================================ def detect_part_colors(texture, num_parts=4, sample_size=200000, color_distance=10): """ 从 segmentation texture 自动发现 Part 颜色。 先统计高频颜色,再 KMeans。 过滤: 黑色 接近白色 极低饱和度背景 """ h, w = texture.shape[:2] pixels = texture.reshape(-1, 3) # -------------------------------------------------------- # 随机采样 # -------------------------------------------------------- if len(pixels) > sample_size: rng = np.random.default_rng(42) idx = rng.choice(len(pixels), sample_size, replace=False) pixels = pixels[idx] pixels_f = pixels.astype(np.float32) # -------------------------------------------------------- # 去掉纯黑/纯白 # -------------------------------------------------------- brightness = pixels_f.mean(axis=1) valid = ((brightness > 5) & (brightness < 250)) pixels_f = pixels_f[valid] print("用于颜色聚类的像素:", len(pixels_f)) # -------------------------------------------------------- # KMeans # -------------------------------------------------------- centers = kmeans_colors(pixels_f, k=num_parts) centers = np.rint(centers).astype(np.uint8) # -------------------------------------------------------- # 按颜色亮度排序,仅为了输出稳定 # -------------------------------------------------------- order = np.argsort(centers.mean(axis=1)) centers = centers[order] print() print("检测到的 Part 颜色:") print("--------------------------------") for i, c in enumerate(centers): print(f"Part {i}: " f"RGB=({c[0]}, {c[1]}, {c[2]})") print("--------------------------------") return centers # ============================================================ # UV 顶点 -> Part # ============================================================ def classify_vertices(texture, uv, centers): """ 给每个 Mesh 顶点确定 Part。 使用: UV -> texture -> 最近颜色中心 """ rgb = sample_texture(texture, uv).astype(np.float32) centers_f = centers.astype(np.float32) dist = ((rgb[:, None, :] - centers_f[None, :, :]) ** 2).sum(axis=2) labels = np.argmin(dist, axis=1) return labels # ============================================================ # Face Part Label # ============================================================ def classify_faces(faces, vertex_labels): """ 三角形的 Part Label。 不简单取第一个顶点。 使用三个顶点投票: [0,0,0] -> 0 [0,1,0] -> 0 [1,1,2] -> 1 这样可以降低边界误判。 """ labels = vertex_labels[faces] face_labels = np.zeros(len(faces), dtype=np.int32) for i in range(len(faces)): a, b, c = labels[i] if a == b or a == c: face_labels[i] = a elif b == c: face_labels[i] = b else: # 三个顶点都不同 # 选择第一个 face_labels[i] = a return face_labels # ============================================================ # 根据 Face Label 提取 Mesh # ============================================================ def extract_part_mesh(mesh, face_indices): """ 从原 Mesh 中提取指定三角形。 """ face_indices = np.asarray(face_indices, dtype=np.int64) faces = mesh.faces[face_indices] # -------------------------------------------------------- # 找到实际使用的顶点 # -------------------------------------------------------- unique_vertices, inverse = np.unique(faces.reshape(-1), return_inverse=True) vertices = mesh.vertices[unique_vertices] new_faces = inverse.reshape(-1, 3) part = trimesh.Trimesh(vertices=vertices, faces=new_faces, process=False) # part.remove_duplicate_faces() # part.remove_degenerate_faces() # part.remove_unreferenced_vertices() # 清理重复面 part.update_faces(part.unique_faces()) # 清理退化三角形 part.update_faces(part.nondegenerate_faces()) # 删除未使用顶点 part.remove_unreferenced_vertices() return part # ============================================================ # 删除小连通域 # ============================================================ def split_components(mesh, min_faces=100): """ 将同一个 Part 中不连接的区域继续拆开。 例如: Part 0 ├── body ├── wheel └── wheel 会得到三个 mesh。 """ try: components = mesh.split(only_watertight=False) except Exception: return [mesh] results = [] for comp in components: if len(comp.faces) < min_faces: continue comp.remove_unreferenced_vertices() results.append(comp) # 如果 split 后什么都没有 if not results and len(mesh.faces) > 0: results.append(mesh) return results # ============================================================ # 创建 Part 材质 # ============================================================ def assign_part_material(mesh, color): """ 给拆出来的 Mesh 使用纯色。 这样导出的 GLB 不再依赖原 segmentation texture。 """ rgba = np.array([int(color[0]), int(color[1]), int(color[2]), 255], dtype=np.uint8) mesh.visual = trimesh.visual.ColorVisuals(mesh=mesh, face_colors=np.tile(rgba, (len(mesh.faces), 1))) return mesh # ============================================================ # 主函数 # ============================================================ def split_segvigen_glb(input_glb, output_dir, num_parts=4, min_faces=100): os.makedirs(output_dir, exist_ok=True) print() print("=" * 70) print("SegviGen GLB Part Splitter") print("=" * 70) scene = trimesh.load(input_glb, force="scene", process=False) print("Geometry:", len(scene.geometry)) # ======================================================== # 2. 提取 Texture # ======================================================== print() print("[2/7] Extracting textures...") images = extract_glb_images(input_glb) print("Embedded images:", len(images)) for i, im in enumerate(images): print(f" image[{i}]: " f"{im.width}x{im.height}") texture = find_segmentation_texture(images) # ======================================================== # 3. 合并 Geometry # ======================================================== print() print("[3/7] Preparing mesh...") geometries = [] for name, geom in scene.geometry.items(): if not isinstance(geom, trimesh.Trimesh): continue geometries.append(geom) if not geometries: raise RuntimeError("GLB 中没有 Mesh") if len(geometries) == 1: mesh = geometries[0] else: mesh = trimesh.util.concatenate(geometries) print("Vertices:", len(mesh.vertices)) print("Faces:", len(mesh.faces)) # ======================================================== # 4. UV # ======================================================== print() print("[4/7] Reading UV...") uv = getattr(mesh.visual, "uv", None) if uv is None: raise RuntimeError("这个 GLB 没有 UV," "无法根据 SegviGen texture 拆分。") uv = np.asarray(uv, dtype=np.float32) print("UV:", uv.shape) # ======================================================== # 5. 自动识别 Part # ======================================================== print() print("[5/7] Detecting Part colors...") centers = detect_part_colors(texture, num_parts=num_parts) # -------------------------------------------------------- # Vertex Label # -------------------------------------------------------- vertex_labels = classify_vertices(texture, uv, centers) # -------------------------------------------------------- # Face Label # -------------------------------------------------------- face_labels = classify_faces(mesh.faces, vertex_labels) # ======================================================== # 6. 提取 Part # ======================================================== print() print("[6/7] Extracting parts...") all_parts = [] global_part_id = 0 for part_id in range(num_parts): face_indices = np.where(face_labels == part_id)[0] print() print(f"Part {part_id}: " f"{len(face_indices)} faces") if len(face_indices) < min_faces: print(" skipped: too few faces") continue part_mesh = extract_part_mesh(mesh, face_indices) # ---------------------------------------------------- # 同一个 Part 按连通域拆分 # ---------------------------------------------------- components = [part_mesh] for component_id, component in enumerate(components): if len(component.faces) < min_faces: continue # ------------------------------------------------ # 材质 # ------------------------------------------------ component = assign_part_material(component, centers[part_id]) filename = (f"part_" f"{global_part_id:03d}" f"_seg{part_id:02d}" f"_component{component_id:02d}" f".glb") output_path = os.path.join(output_dir, filename) component.export(output_path) print(" ->", filename, "vertices=", len(component.vertices), "faces=", len(component.faces)) all_parts.append((global_part_id, part_id, component)) global_part_id += 1 # ======================================================== # 7. 输出合并 GLB # ======================================================== print() print("[7/7] Exporting combined GLB...") if all_parts: scene_out = trimesh.Scene() for global_id, part_id, part in all_parts: scene_out.add_geometry(part, node_name=f"part_{global_id:03d}", geom_name=f"part_{global_id:03d}") combined_path = os.path.join(output_dir, "all_parts.glb") scene_out.export(combined_path) print("Combined:", combined_path) print() print("=" * 70) print("完成!共输出", len(all_parts), "个子部件") print("=" * 70) print() print("输出目录:") print(os.path.abspath(output_dir)) def main(): parser = argparse.ArgumentParser(description=("SegviGen segmentation GLB " "automatic part splitter")) parser.add_argument("--input", default=r"E:\project\3d_label\chai_server\output3.glb", help="SegviGen output GLB") parser.add_argument("-o", "--output", default="output_parts", help="输出目录") parser.add_argument("-k", "--num-parts", type=int, default=4, help="Part 数量,默认 4") parser.add_argument("--min-faces", type=int, default=100, help="最小三角形数量") parser.add_argument("--no-components", action="store_true", help="不进一步拆分几何连通域") args = parser.parse_args() split_segvigen_glb(input_glb=args.input, output_dir=args.output, num_parts=args.num_parts, min_faces=args.min_faces) if __name__ == "__main__": main()