diff --git a/example_workflows/MeshOnly_Pixal3D.json b/example_workflows/MeshOnly_Pixal3D.json new file mode 100644 index 0000000..121f4d7 --- /dev/null +++ b/example_workflows/MeshOnly_Pixal3D.json @@ -0,0 +1,1335 @@ +{ + "id": "e701b663-dd55-4899-a4c2-333fc48d84d3", + "revision": 0, + "last_node_id": 288, + "last_link_id": 564, + "nodes": [ + { + "id": 250, + "type": "Trellis2FillHolesNicelyWithMeshlib", + "pos": [ + -54.768657050258554, + 704.3729247955466 + ], + "size": [ + 328.5106969603829, + 46 + ], + "flags": {}, + "order": 13, + "mode": 4, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 463 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 488 + ] + }, + { + "name": "holes_filled", + "type": "INT", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "9d1f80ec9b4d5c4b966e7595ee99e01a1d0c9f19", + "Node name for S&R": "Trellis2FillHolesNicelyWithMeshlib" + }, + "widgets_values": [] + }, + { + "id": 257, + "type": "Trellis2SimplifyMesh", + "pos": [ + -26.817768803861814, + 816.7701525553716 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 14, + "mode": 4, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 488 + }, + { + "name": "target_face_num", + "type": "INT", + "widget": { + "name": "target_face_num" + }, + "link": 487 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 489 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2SimplifyMesh" + }, + "widgets_values": [ + 500000, + "Cumesh" + ] + }, + { + "id": 6, + "type": "Trellis2LoadImageWithTransparency", + "pos": [ + -1898.8408811895958, + 425.13094800380867 + ], + "size": [ + 631.9393495485403, + 765.9737728955563 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": null + }, + { + "name": "mask", + "type": "MASK", + "links": null + }, + { + "name": "image_with_alpha", + "type": "IMAGE", + "links": [ + 344 + ] + } + ], + "title": "Image with Transparency", + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "e7f9b30df7a09bcedf1c955e176754c73f983254", + "Node name for S&R": "Trellis2LoadImageWithTransparency", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "DwarfSorcerer.png", + "image" + ] + }, + { + "id": 194, + "type": "Trellis2PreProcessImage", + "pos": [ + -1218.495705914309, + 437.03938555658084 + ], + "size": [ + 281.7837890625, + 106 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 344 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 394, + 515, + 522, + 532 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "7ce26deae425d114f3deb78e802d6196e5987a4a", + "Node name for S&R": "Trellis2PreProcessImage" + }, + "widgets_values": [ + 5, + false, + 1024 + ] + }, + { + "id": 213, + "type": "Trellis2ImageCondGenerator", + "pos": [ + -873.6756939901774, + 155.23142744547695 + ], + "size": [ + 308.0748046875, + 118 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 395 + }, + { + "name": "image", + "type": "IMAGE", + "link": 394 + } + ], + "outputs": [ + { + "name": "cond_512", + "type": "IMAGE_COND", + "links": [ + 521, + 526 + ] + }, + { + "name": "cond_1024", + "type": "IMAGE_COND", + "links": [ + 530 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 520 + ] + }, + { + "name": "moge_camera_config", + "type": "MOGE_CAM_CONFIG", + "links": [ + 516, + 523, + 533 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2ImageCondGenerator" + }, + "widgets_values": [ + 1 + ] + }, + { + "id": 217, + "type": "Trellis2DecodeLatents", + "pos": [ + -371.3942860879976, + 839.9213276329912 + ], + "size": [ + 210, + 122 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 537 + }, + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "link": 536 + }, + { + "name": "texture_slat", + "shape": 7, + "type": "TEXTURE_SLAT", + "link": null + }, + { + "name": "resolution", + "type": "INT", + "widget": { + "name": "resolution" + }, + "link": 538 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 397 + ] + }, + { + "name": "bvh", + "type": "BVH", + "links": [] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2DecodeLatents" + }, + "widgets_values": [ + 0, + true + ] + }, + { + "id": 193, + "type": "Trellis2FillHolesWithCuMesh", + "pos": [ + -23.408733901895896, + 210.02490495474518 + ], + "size": [ + 312.4361328125, + 58 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 397 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 356 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "dc66f263e832b103a81a021bf6cd53eec8ff28af", + "Node name for S&R": "Trellis2FillHolesWithCuMesh" + }, + "widgets_values": [ + 1 + ] + }, + { + "id": 161, + "type": "Trellis2ReconstructMeshWithQuad", + "pos": [ + -39.819037021447066, + 335.3167693354984 + ], + "size": [ + 331.5878996659427, + 142.11360851901668 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 356 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 398 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "d4d66cfc9f1ce5b1c3b9a2e504ee4b24ffad48a5", + "Node name for S&R": "Trellis2ReconstructMeshWithQuad", + "widget_ue_connectable": {} + }, + "widgets_values": [ + 1, + 1024, + true, + true + ] + }, + { + "id": 218, + "type": "Trellis2SimplifyMesh", + "pos": [ + 8.858448178748294, + 541.0390332148819 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 398 + }, + { + "name": "target_face_num", + "type": "INT", + "widget": { + "name": "target_face_num" + }, + "link": 454 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 463 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2SimplifyMesh" + }, + "widgets_values": [ + 500000, + "Cumesh" + ] + }, + { + "id": 209, + "type": "PrimitiveInt", + "pos": [ + -1830.8929517313256, + 263.9724281177611 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 454, + 487 + ] + } + ], + "title": "Target Face Number", + "properties": { + "cnr_id": "comfy-core", + "ver": "0.16.4", + "Node name for S&R": "PrimitiveInt" + }, + "widgets_values": [ + 1000000, + "fixed" + ] + }, + { + "id": 202, + "type": "Trellis2MeshWithVoxelToTrimesh", + "pos": [ + -40.07205166978725, + 967.0539753408918 + ], + "size": [ + 349.41171875, + 58 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 489 + } + ], + "outputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "links": [ + 456 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2MeshWithVoxelToTrimesh" + }, + "widgets_values": [ + "None" + ] + }, + { + "id": 203, + "type": "Trellis2ExportMesh", + "pos": [ + -21.46664432202061, + 1121.318173070784 + ], + "size": [ + 270, + 102 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "link": 456 + }, + { + "name": "filename_prefix", + "type": "STRING", + "widget": { + "name": "filename_prefix" + }, + "link": 402 + } + ], + "outputs": [ + { + "name": "glb_path", + "type": "STRING", + "links": [ + 429 + ] + }, + { + "name": "relative_path", + "type": "STRING", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2ExportMesh" + }, + "widgets_values": [ + "LaserPistol", + "glb" + ] + }, + { + "id": 232, + "type": "Preview3D", + "pos": [ + 364.38572582972074, + 83.54091160006845 + ], + "size": [ + 1197.5811272300746, + 1196.1310737028766 + ], + "flags": {}, + "order": 17, + "mode": 0, + "inputs": [ + { + "name": "camera_info", + "shape": 7, + "type": "LOAD3D_CAMERA", + "link": null + }, + { + "name": "bg_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "model_file", + "type": "STRING,FILE_3D_GLB,FILE_3D_GLTF,FILE_3D_FBX,FILE_3D_OBJ,FILE_3D_STL,FILE_3D_USDZ,FILE_3D", + "widget": { + "name": "model_file" + }, + "link": 429 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.16.4", + "Node name for S&R": "Preview3D", + "Last Time Model File": "C:/Git/ComfyUI/output/DwarfSorcerer_Trellis2_00001_.glb", + "Resource Folder": "Git/ComfyUI/output", + "Scene Config": { + "showGrid": true, + "backgroundColor": "#282828", + "backgroundImage": "", + "backgroundRenderMode": "tiled" + }, + "Camera Config": { + "cameraType": "perspective", + "fov": 35, + "state": { + "position": { + "x": -0.0020336302359329943, + "y": -9.086858426343413, + "z": 0.8590716845817458 + }, + "target": { + "x": 0, + "y": 0.9516711391676895, + "z": 0 + }, + "zoom": 1, + "cameraType": "perspective" + } + }, + "Light Config": { + "intensity": 3 + }, + "Model Config": { + "upDirection": "original", + "materialMode": "normal", + "showSkeleton": false + } + }, + "widgets_values": [ + "C:/Git/ComfyUI/output/DwarfSorcerer_Trellis2_00001_.glb", + "" + ] + }, + { + "id": 272, + "type": "Trellis2SparseGenerator", + "pos": [ + -874.3328359531824, + 351.98478429322415 + ], + "size": [ + 395.65625, + 526 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 520 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 521 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 522 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 523 + } + ], + "outputs": [ + { + "name": "coords", + "type": "COORDS", + "links": [ + 524 + ] + }, + { + "name": "sparse_structure_resolution", + "type": "INT", + "links": [ + 535 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 525 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "86d13d9eac4a2bd4395954c7184d3aa4fa81a9d8", + "Node name for S&R": "Trellis2SparseGenerator" + }, + "widgets_values": [ + 12345, + "fixed", + 25, + 7.5, + 0.01, + 5, + "heun", + 32, + 0.1, + 1, + true, + 1, + false, + 0.65, + 2, + "flood_fill", + 0.92, + true + ] + }, + { + "id": 215, + "type": "Trellis2ShapeGenerator", + "pos": [ + -841.1718308975657, + 945.0367148651351 + ], + "size": [ + 335.24609375, + 402 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 525 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 526 + }, + { + "name": "coords", + "type": "COORDS", + "link": 524 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 515 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 516 + } + ], + "outputs": [ + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "links": [ + 531 + ] + }, + { + "name": "resolution", + "type": "INT", + "links": [ + 534 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 529 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2ShapeGenerator" + }, + "widgets_values": [ + 512, + 25, + 7.5, + 0.01, + 3, + "euler", + 0.1, + 1, + false, + 0.65, + 2, + 0.92 + ] + }, + { + "id": 274, + "type": "Trellis2ShapeCascadeGenerator", + "pos": [ + -432.74553813924416, + 316.7368953282102 + ], + "size": [ + 335.9791015625, + 474 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 529 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 530 + }, + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "link": 531 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 532 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 533 + }, + { + "name": "from_resolution", + "type": "INT", + "widget": { + "name": "from_resolution" + }, + "link": 534 + }, + { + "name": "sparse_structure_resolution", + "type": "INT", + "widget": { + "name": "sparse_structure_resolution" + }, + "link": 535 + } + ], + "outputs": [ + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "links": [ + 536 + ] + }, + { + "name": "resolution", + "type": "INT", + "links": [ + 538 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 537 + ] + }, + { + "name": "num_tokens", + "type": "INT", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "86d13d9eac4a2bd4395954c7184d3aa4fa81a9d8", + "Node name for S&R": "Trellis2ShapeCascadeGenerator" + }, + "widgets_values": [ + 0, + 1024, + 32, + 999999, + 25, + 7.5, + 0.01, + 3, + "heun", + 0.1, + 1, + false, + 0.65, + 2, + 0.92 + ] + }, + { + "id": 39, + "type": "Trellis2LoadModel", + "pos": [ + -1230.97479805832, + 117.1260927127055 + ], + "size": [ + 303.21076796905936, + 226 + ], + "flags": { + "collapsed": false + }, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 395 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "c5d66aa46bf1bbf903fd4fe9ff086ac774bb7e03", + "Node name for S&R": "Trellis2LoadModel", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "TencentARC/Pixal3D-T", + "flash_attn", + "cuda", + true, + false, + "flex_gemm", + "flash_attn", + false + ] + }, + { + "id": 219, + "type": "PrimitiveString", + "pos": [ + -1830.5311181322465, + 133.64378105828288 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 402 + ] + } + ], + "title": "Name", + "properties": { + "cnr_id": "comfy-core", + "ver": "0.16.4", + "Node name for S&R": "PrimitiveString" + }, + "widgets_values": [ + "Pixal3D" + ] + } + ], + "links": [ + [ + 344, + 6, + 2, + 194, + 0, + "IMAGE" + ], + [ + 356, + 193, + 0, + 161, + 0, + "MESHWITHVOXEL" + ], + [ + 394, + 194, + 0, + 213, + 1, + "IMAGE" + ], + [ + 395, + 39, + 0, + 213, + 0, + "TRELLIS2PIPELINE" + ], + [ + 397, + 217, + 0, + 193, + 0, + "MESHWITHVOXEL" + ], + [ + 398, + 161, + 0, + 218, + 0, + "MESHWITHVOXEL" + ], + [ + 402, + 219, + 0, + 203, + 1, + "STRING" + ], + [ + 429, + 203, + 0, + 232, + 2, + "STRING" + ], + [ + 454, + 209, + 0, + 218, + 1, + "INT" + ], + [ + 456, + 202, + 0, + 203, + 0, + "TRIMESH" + ], + [ + 463, + 218, + 0, + 250, + 0, + "MESHWITHVOXEL" + ], + [ + 487, + 209, + 0, + 257, + 1, + "INT" + ], + [ + 488, + 250, + 0, + 257, + 0, + "MESHWITHVOXEL" + ], + [ + 489, + 257, + 0, + 202, + 0, + "MESHWITHVOXEL" + ], + [ + 515, + 194, + 0, + 215, + 3, + "IMAGE" + ], + [ + 516, + 213, + 3, + 215, + 4, + "MOGE_CAM_CONFIG" + ], + [ + 520, + 213, + 2, + 272, + 0, + "TRELLIS2PIPELINE" + ], + [ + 521, + 213, + 0, + 272, + 1, + "IMAGE_COND" + ], + [ + 522, + 194, + 0, + 272, + 2, + "IMAGE" + ], + [ + 523, + 213, + 3, + 272, + 3, + "MOGE_CAM_CONFIG" + ], + [ + 524, + 272, + 0, + 215, + 2, + "COORDS" + ], + [ + 525, + 272, + 2, + 215, + 0, + "TRELLIS2PIPELINE" + ], + [ + 526, + 213, + 0, + 215, + 1, + "IMAGE_COND" + ], + [ + 529, + 215, + 2, + 274, + 0, + "TRELLIS2PIPELINE" + ], + [ + 530, + 213, + 1, + 274, + 1, + "IMAGE_COND" + ], + [ + 531, + 215, + 0, + 274, + 2, + "SHAPE_SLAT" + ], + [ + 532, + 194, + 0, + 274, + 3, + "IMAGE" + ], + [ + 533, + 213, + 3, + 274, + 4, + "MOGE_CAM_CONFIG" + ], + [ + 534, + 215, + 1, + 274, + 5, + "INT" + ], + [ + 535, + 272, + 1, + 274, + 6, + "INT" + ], + [ + 536, + 274, + 0, + 217, + 1, + "SHAPE_SLAT" + ], + [ + 537, + 274, + 2, + 217, + 0, + "TRELLIS2PIPELINE" + ], + [ + 538, + 274, + 1, + 217, + 3, + "INT" + ] + ], + "groups": [ + { + "id": 2, + "title": "Configuration", + "bounding": [ + -1908.8408811895958, + 51.87071308817088, + 651.9393495485403, + 1149.234007811194 + ], + "color": "#3f789e", + "flags": {} + }, + { + "id": 3, + "title": "Generation", + "bounding": [ + -1240.97479805832, + 43.52609271270552, + 1584.4033789519756, + 1342.8512046673666 + ], + "color": "#3f789e", + "flags": {} + } + ], + "config": {}, + "extra": { + "workflowRendererVersion": "LG", + "ue_links": [], + "ds": { + "scale": 0.4718841099024524, + "offset": [ + 2039.8780143037789, + 220.6100941969397 + ] + }, + "links_added_by_ue": [], + "frontendVersion": "1.43.18", + "VHS_latentpreview": false, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/example_workflows/MeshWithTexturing_Pixal3D.json b/example_workflows/MeshWithTexturing_Pixal3D.json new file mode 100644 index 0000000..c41a553 --- /dev/null +++ b/example_workflows/MeshWithTexturing_Pixal3D.json @@ -0,0 +1,1501 @@ +{ + "id": "0cc31dc7-54c1-494d-b20e-d9f4aabc66bc", + "revision": 0, + "last_node_id": 290, + "last_link_id": 574, + "nodes": [ + { + "id": 6, + "type": "Trellis2LoadImageWithTransparency", + "pos": [ + -1898.8408811895958, + 425.13094800380867 + ], + "size": [ + 631.9393495485403, + 765.9737728955563 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": null + }, + { + "name": "mask", + "type": "MASK", + "links": null + }, + { + "name": "image_with_alpha", + "type": "IMAGE", + "links": [ + 344 + ] + } + ], + "title": "Image with Transparency", + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "e7f9b30df7a09bcedf1c955e176754c73f983254", + "Node name for S&R": "Trellis2LoadImageWithTransparency", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "DwarfSorcerer.png", + "image" + ] + }, + { + "id": 209, + "type": "PrimitiveInt", + "pos": [ + -1830.8929517313256, + 263.9724281177611 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 454, + 487 + ] + } + ], + "title": "Target Face Number", + "properties": { + "cnr_id": "comfy-core", + "ver": "0.16.4", + "Node name for S&R": "PrimitiveInt" + }, + "widgets_values": [ + 1000000, + "fixed" + ] + }, + { + "id": 272, + "type": "Trellis2SparseGenerator", + "pos": [ + -874.3328359531824, + 351.98478429322415 + ], + "size": [ + 395.65625, + 526 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 520 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 521 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 522 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 523 + } + ], + "outputs": [ + { + "name": "coords", + "type": "COORDS", + "links": [ + 524 + ] + }, + { + "name": "sparse_structure_resolution", + "type": "INT", + "links": [ + 535 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 525 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "86d13d9eac4a2bd4395954c7184d3aa4fa81a9d8", + "Node name for S&R": "Trellis2SparseGenerator" + }, + "widgets_values": [ + 12345, + "fixed", + 25, + 7.5, + 0.01, + 5, + "heun", + 32, + 0.1, + 1, + true, + 1, + false, + 0.65, + 2, + "flood_fill", + 0.92, + true + ] + }, + { + "id": 215, + "type": "Trellis2ShapeGenerator", + "pos": [ + -841.1718308975657, + 945.0367148651351 + ], + "size": [ + 335.24609375, + 402 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 525 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 526 + }, + { + "name": "coords", + "type": "COORDS", + "link": 524 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 515 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 516 + } + ], + "outputs": [ + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "links": [ + 531 + ] + }, + { + "name": "resolution", + "type": "INT", + "links": [ + 534 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 529 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2ShapeGenerator" + }, + "widgets_values": [ + 512, + 25, + 7.5, + 0.01, + 3, + "euler", + 0.1, + 1, + false, + 0.65, + 2, + 0.92 + ] + }, + { + "id": 39, + "type": "Trellis2LoadModel", + "pos": [ + -1230.97479805832, + 117.1260927127055 + ], + "size": [ + 303.21076796905936, + 226 + ], + "flags": { + "collapsed": false + }, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 395 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "c5d66aa46bf1bbf903fd4fe9ff086ac774bb7e03", + "Node name for S&R": "Trellis2LoadModel", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "TencentARC/Pixal3D-T", + "flash_attn", + "cuda", + true, + false, + "flex_gemm", + "flash_attn", + false + ] + }, + { + "id": 274, + "type": "Trellis2ShapeCascadeGenerator", + "pos": [ + -420.008261589595, + 174.50398195526927 + ], + "size": [ + 335.9791015625, + 474 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 529 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 530 + }, + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "link": 531 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 532 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 533 + }, + { + "name": "from_resolution", + "type": "INT", + "widget": { + "name": "from_resolution" + }, + "link": 534 + }, + { + "name": "sparse_structure_resolution", + "type": "INT", + "widget": { + "name": "sparse_structure_resolution" + }, + "link": 535 + } + ], + "outputs": [ + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "links": [ + 536, + 567 + ] + }, + { + "name": "resolution", + "type": "INT", + "links": [ + 538 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 565 + ] + }, + { + "name": "num_tokens", + "type": "INT", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "86d13d9eac4a2bd4395954c7184d3aa4fa81a9d8", + "Node name for S&R": "Trellis2ShapeCascadeGenerator" + }, + "widgets_values": [ + 0, + 1024, + 32, + 999999, + 25, + 7.5, + 0.01, + 3, + "heun", + 0.1, + 1, + false, + 0.65, + 2, + 0.92 + ] + }, + { + "id": 194, + "type": "Trellis2PreProcessImage", + "pos": [ + -1218.495705914309, + 437.03938555658084 + ], + "size": [ + 281.7837890625, + 106 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 344 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 394, + 515, + 522, + 532, + 568 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "7ce26deae425d114f3deb78e802d6196e5987a4a", + "Node name for S&R": "Trellis2PreProcessImage" + }, + "widgets_values": [ + 5, + false, + 1024 + ] + }, + { + "id": 213, + "type": "Trellis2ImageCondGenerator", + "pos": [ + -873.6756939901774, + 155.23142744547695 + ], + "size": [ + 308.0748046875, + 118 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 395 + }, + { + "name": "image", + "type": "IMAGE", + "link": 394 + } + ], + "outputs": [ + { + "name": "cond_512", + "type": "IMAGE_COND", + "links": [ + 521, + 526 + ] + }, + { + "name": "cond_1024", + "type": "IMAGE_COND", + "links": [ + 530, + 566 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 520 + ] + }, + { + "name": "moge_camera_config", + "type": "MOGE_CAM_CONFIG", + "links": [ + 516, + 523, + 533, + 569 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2ImageCondGenerator" + }, + "widgets_values": [ + 1 + ] + }, + { + "id": 289, + "type": "Trellis2TexSlatGenerator", + "pos": [ + -410.72426621796166, + 724.3885535781652 + ], + "size": [ + 340.390625, + 402 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 565 + }, + { + "name": "image_cond", + "type": "IMAGE_COND", + "link": 566 + }, + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "link": 567 + }, + { + "name": "image", + "shape": 7, + "type": "IMAGE", + "link": 568 + }, + { + "name": "moge_camera_config", + "shape": 7, + "type": "MOGE_CAM_CONFIG", + "link": 569 + } + ], + "outputs": [ + { + "name": "texture_slat", + "type": "TEXTURE_SLAT", + "links": [ + 570 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 571 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "86d13d9eac4a2bd4395954c7184d3aa4fa81a9d8", + "Node name for S&R": "Trellis2TexSlatGenerator" + }, + "widgets_values": [ + 1024, + 25, + 7.5, + 0, + 3, + "euler", + 0, + 0.9, + false, + 0.65, + 2, + 0.92 + ] + }, + { + "id": 193, + "type": "Trellis2FillHolesWithCuMesh", + "pos": [ + 184.40535063876197, + 149.86491249810132 + ], + "size": [ + 312.4361328125, + 58 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 397 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 356 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "dc66f263e832b103a81a021bf6cd53eec8ff28af", + "Node name for S&R": "Trellis2FillHolesWithCuMesh" + }, + "widgets_values": [ + 1 + ] + }, + { + "id": 161, + "type": "Trellis2ReconstructMeshWithQuad", + "pos": [ + 189.04847750020792, + 267.26177075303883 + ], + "size": [ + 331.5878996659427, + 142.11360851901668 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 356 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 398 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "d4d66cfc9f1ce5b1c3b9a2e504ee4b24ffad48a5", + "Node name for S&R": "Trellis2ReconstructMeshWithQuad", + "widget_ue_connectable": {} + }, + "widgets_values": [ + 1, + 1024, + true, + true + ] + }, + { + "id": 218, + "type": "Trellis2SimplifyMesh", + "pos": [ + 219.30428174016674, + 473.8612441499314 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 398 + }, + { + "name": "target_face_num", + "type": "INT", + "widget": { + "name": "target_face_num" + }, + "link": 454 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 463 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2SimplifyMesh" + }, + "widgets_values": [ + 500000, + "Cumesh" + ] + }, + { + "id": 250, + "type": "Trellis2FillHolesNicelyWithMeshlib", + "pos": [ + 198.6610452102749, + 613.5101575092269 + ], + "size": [ + 328.5106969603829, + 46 + ], + "flags": {}, + "order": 14, + "mode": 4, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 463 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 488 + ] + }, + { + "name": "holes_filled", + "type": "INT", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "9d1f80ec9b4d5c4b966e7595ee99e01a1d0c9f19", + "Node name for S&R": "Trellis2FillHolesNicelyWithMeshlib" + }, + "widgets_values": [] + }, + { + "id": 203, + "type": "Trellis2ExportMesh", + "pos": [ + 170.5575082791613, + 1283.095561654488 + ], + "size": [ + 270, + 102 + ], + "flags": {}, + "order": 17, + "mode": 0, + "inputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "link": 574 + }, + { + "name": "filename_prefix", + "type": "STRING", + "widget": { + "name": "filename_prefix" + }, + "link": 402 + } + ], + "outputs": [ + { + "name": "glb_path", + "type": "STRING", + "links": [ + 429 + ] + }, + { + "name": "relative_path", + "type": "STRING", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2ExportMesh" + }, + "widgets_values": [ + "LaserPistol", + "glb" + ] + }, + { + "id": 257, + "type": "Trellis2SimplifyMesh", + "pos": [ + 222.22584571304807, + 705.7312451177193 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 15, + "mode": 4, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 488 + }, + { + "name": "target_face_num", + "type": "INT", + "widget": { + "name": "target_face_num" + }, + "link": 487 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 572 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2SimplifyMesh" + }, + "widgets_values": [ + 500000, + "Cumesh" + ] + }, + { + "id": 217, + "type": "Trellis2DecodeLatents", + "pos": [ + -48.8569261943387, + 548.0253905072441 + ], + "size": [ + 210, + 122 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 571 + }, + { + "name": "shape_slat", + "type": "SHAPE_SLAT", + "link": 536 + }, + { + "name": "texture_slat", + "shape": 7, + "type": "TEXTURE_SLAT", + "link": 570 + }, + { + "name": "resolution", + "type": "INT", + "widget": { + "name": "resolution" + }, + "link": 538 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 397 + ] + }, + { + "name": "bvh", + "type": "BVH", + "links": [ + 573 + ] + }, + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "ef2286b47ce0ef0e681dffe0209aadd96de50f72", + "Node name for S&R": "Trellis2DecodeLatents" + }, + "widgets_values": [ + 0, + true + ] + }, + { + "id": 290, + "type": "Trellis2UnWrapAndRasterizer", + "pos": [ + 104.1595211548057, + 859.3392446479718 + ], + "size": [ + 419.15234375, + 338 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 572 + }, + { + "name": "bvh", + "type": "BVH", + "link": 573 + } + ], + "outputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "links": [ + 574 + ] + }, + { + "name": "base_color_texture", + "type": "IMAGE", + "links": null + }, + { + "name": "metallic_roughness_texture", + "type": "IMAGE", + "links": null + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "86d13d9eac4a2bd4395954c7184d3aa4fa81a9d8", + "Node name for S&R": "Trellis2UnWrapAndRasterizer" + }, + "widgets_values": [ + 60, + 0, + 1, + 1, + 2048, + "OPAQUE", + false, + false, + false, + "telea", + "None" + ] + }, + { + "id": 232, + "type": "Preview3D", + "pos": [ + 584.4700506848733, + 118.27654774792936 + ], + "size": [ + 1197.5811272300746, + 1196.1310737028766 + ], + "flags": {}, + "order": 18, + "mode": 0, + "inputs": [ + { + "name": "camera_info", + "shape": 7, + "type": "LOAD3D_CAMERA", + "link": null + }, + { + "name": "bg_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "model_file", + "type": "STRING,FILE_3D_GLB,FILE_3D_GLTF,FILE_3D_FBX,FILE_3D_OBJ,FILE_3D_STL,FILE_3D_USDZ,FILE_3D", + "widget": { + "name": "model_file" + }, + "link": 429 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.16.4", + "Node name for S&R": "Preview3D", + "Last Time Model File": "C:/Git/ComfyUI/output/Pixal3D_Textured_00001_.glb", + "Resource Folder": "Git/ComfyUI/output", + "Scene Config": { + "showGrid": true, + "backgroundColor": "#282828", + "backgroundImage": "", + "backgroundRenderMode": "tiled" + }, + "Camera Config": { + "cameraType": "perspective", + "fov": 35, + "state": { + "position": { + "x": 0.2051689293676939, + "y": 4.275591363360545, + "z": 9.50787749809354 + }, + "target": { + "x": 0, + "y": 2.5, + "z": 0 + }, + "zoom": 1, + "cameraType": "perspective" + } + }, + "Light Config": { + "intensity": 3 + }, + "Model Config": { + "upDirection": "original", + "materialMode": "original", + "showSkeleton": false + } + }, + "widgets_values": [ + "C:/Git/ComfyUI/output/Pixal3D_Textured_00001_.glb", + "" + ] + }, + { + "id": 219, + "type": "PrimitiveString", + "pos": [ + -1830.5311181322465, + 133.64378105828288 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 402 + ] + } + ], + "title": "Name", + "properties": { + "cnr_id": "comfy-core", + "ver": "0.16.4", + "Node name for S&R": "PrimitiveString" + }, + "widgets_values": [ + "Pixal3D_Textured" + ] + } + ], + "links": [ + [ + 344, + 6, + 2, + 194, + 0, + "IMAGE" + ], + [ + 356, + 193, + 0, + 161, + 0, + "MESHWITHVOXEL" + ], + [ + 394, + 194, + 0, + 213, + 1, + "IMAGE" + ], + [ + 395, + 39, + 0, + 213, + 0, + "TRELLIS2PIPELINE" + ], + [ + 397, + 217, + 0, + 193, + 0, + "MESHWITHVOXEL" + ], + [ + 398, + 161, + 0, + 218, + 0, + "MESHWITHVOXEL" + ], + [ + 402, + 219, + 0, + 203, + 1, + "STRING" + ], + [ + 429, + 203, + 0, + 232, + 2, + "STRING" + ], + [ + 454, + 209, + 0, + 218, + 1, + "INT" + ], + [ + 463, + 218, + 0, + 250, + 0, + "MESHWITHVOXEL" + ], + [ + 487, + 209, + 0, + 257, + 1, + "INT" + ], + [ + 488, + 250, + 0, + 257, + 0, + "MESHWITHVOXEL" + ], + [ + 515, + 194, + 0, + 215, + 3, + "IMAGE" + ], + [ + 516, + 213, + 3, + 215, + 4, + "MOGE_CAM_CONFIG" + ], + [ + 520, + 213, + 2, + 272, + 0, + "TRELLIS2PIPELINE" + ], + [ + 521, + 213, + 0, + 272, + 1, + "IMAGE_COND" + ], + [ + 522, + 194, + 0, + 272, + 2, + "IMAGE" + ], + [ + 523, + 213, + 3, + 272, + 3, + "MOGE_CAM_CONFIG" + ], + [ + 524, + 272, + 0, + 215, + 2, + "COORDS" + ], + [ + 525, + 272, + 2, + 215, + 0, + "TRELLIS2PIPELINE" + ], + [ + 526, + 213, + 0, + 215, + 1, + "IMAGE_COND" + ], + [ + 529, + 215, + 2, + 274, + 0, + "TRELLIS2PIPELINE" + ], + [ + 530, + 213, + 1, + 274, + 1, + "IMAGE_COND" + ], + [ + 531, + 215, + 0, + 274, + 2, + "SHAPE_SLAT" + ], + [ + 532, + 194, + 0, + 274, + 3, + "IMAGE" + ], + [ + 533, + 213, + 3, + 274, + 4, + "MOGE_CAM_CONFIG" + ], + [ + 534, + 215, + 1, + 274, + 5, + "INT" + ], + [ + 535, + 272, + 1, + 274, + 6, + "INT" + ], + [ + 536, + 274, + 0, + 217, + 1, + "SHAPE_SLAT" + ], + [ + 538, + 274, + 1, + 217, + 3, + "INT" + ], + [ + 565, + 274, + 2, + 289, + 0, + "TRELLIS2PIPELINE" + ], + [ + 566, + 213, + 1, + 289, + 1, + "IMAGE_COND" + ], + [ + 567, + 274, + 0, + 289, + 2, + "SHAPE_SLAT" + ], + [ + 568, + 194, + 0, + 289, + 3, + "IMAGE" + ], + [ + 569, + 213, + 3, + 289, + 4, + "MOGE_CAM_CONFIG" + ], + [ + 570, + 289, + 0, + 217, + 2, + "TEXTURE_SLAT" + ], + [ + 571, + 289, + 1, + 217, + 0, + "TRELLIS2PIPELINE" + ], + [ + 572, + 257, + 0, + 290, + 0, + "MESHWITHVOXEL" + ], + [ + 573, + 217, + 1, + 290, + 1, + "BVH" + ], + [ + 574, + 290, + 0, + 203, + 0, + "TRIMESH" + ] + ], + "groups": [ + { + "id": 2, + "title": "Configuration", + "bounding": [ + -1908.8408811895958, + 51.87071308817088, + 651.9393495485403, + 1149.234007811194 + ], + "color": "#3f789e", + "flags": {} + }, + { + "id": 3, + "title": "Generation", + "bounding": [ + -1240.97479805832, + 43.52609271270552, + 1788.199900924072, + 1349.2198429421912 + ], + "color": "#3f789e", + "flags": {} + } + ], + "config": {}, + "extra": { + "workflowRendererVersion": "LG", + "ue_links": [], + "ds": { + "scale": 0.4718841099024524, + "offset": [ + 1731.2302331480612, + 289.8930889868392 + ] + }, + "links_added_by_ue": [], + "frontendVersion": "1.43.18", + "VHS_latentpreview": false, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/moge/__init__.py b/moge/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/moge/model/__init__.py b/moge/model/__init__.py new file mode 100644 index 0000000..c919e3b --- /dev/null +++ b/moge/model/__init__.py @@ -0,0 +1,18 @@ +import importlib +from typing import * + +if TYPE_CHECKING: + from .v1 import MoGeModel as MoGeModelV1 + from .v2 import MoGeModel as MoGeModelV2 + + +def import_model_class_by_version(version: str) -> Type[Union['MoGeModelV1', 'MoGeModelV2']]: + assert version in ['v1', 'v2'], f'Unsupported model version: {version}' + + try: + module = importlib.import_module(f'.{version}', __package__) + except ModuleNotFoundError: + raise ValueError(f'Model version "{version}" not found.') + + cls = getattr(module, 'MoGeModel') + return cls diff --git a/moge/model/dinov2/__init__.py b/moge/model/dinov2/__init__.py new file mode 100644 index 0000000..ae847e4 --- /dev/null +++ b/moge/model/dinov2/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +__version__ = "0.0.1" diff --git a/moge/model/dinov2/hub/__init__.py b/moge/model/dinov2/hub/__init__.py new file mode 100644 index 0000000..b88da6b --- /dev/null +++ b/moge/model/dinov2/hub/__init__.py @@ -0,0 +1,4 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. diff --git a/moge/model/dinov2/hub/backbones.py b/moge/model/dinov2/hub/backbones.py new file mode 100644 index 0000000..53fe837 --- /dev/null +++ b/moge/model/dinov2/hub/backbones.py @@ -0,0 +1,156 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +from enum import Enum +from typing import Union + +import torch + +from .utils import _DINOV2_BASE_URL, _make_dinov2_model_name + + +class Weights(Enum): + LVD142M = "LVD142M" + + +def _make_dinov2_model( + *, + arch_name: str = "vit_large", + img_size: int = 518, + patch_size: int = 14, + init_values: float = 1.0, + ffn_layer: str = "mlp", + block_chunks: int = 0, + num_register_tokens: int = 0, + interpolate_antialias: bool = False, + interpolate_offset: float = 0.1, + pretrained: bool = True, + weights: Union[Weights, str] = Weights.LVD142M, + **kwargs, +): + from ..models import vision_transformer as vits + + if isinstance(weights, str): + try: + weights = Weights[weights] + except KeyError: + raise AssertionError(f"Unsupported weights: {weights}") + + model_base_name = _make_dinov2_model_name(arch_name, patch_size) + vit_kwargs = dict( + img_size=img_size, + patch_size=patch_size, + init_values=init_values, + ffn_layer=ffn_layer, + block_chunks=block_chunks, + num_register_tokens=num_register_tokens, + interpolate_antialias=interpolate_antialias, + interpolate_offset=interpolate_offset, + ) + vit_kwargs.update(**kwargs) + model = vits.__dict__[arch_name](**vit_kwargs) + + if pretrained: + model_full_name = _make_dinov2_model_name(arch_name, patch_size, num_register_tokens) + url = _DINOV2_BASE_URL + f"/{model_base_name}/{model_full_name}_pretrain.pth" + state_dict = torch.hub.load_state_dict_from_url(url, map_location="cpu") + model.load_state_dict(state_dict, strict=True) + + return model + + +def dinov2_vits14(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-S/14 model (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model(arch_name="vit_small", pretrained=pretrained, weights=weights, **kwargs) + + +def dinov2_vitb14(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-B/14 model (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model(arch_name="vit_base", pretrained=pretrained, weights=weights, **kwargs) + + +def dinov2_vitl14(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-L/14 model (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model(arch_name="vit_large", pretrained=pretrained, weights=weights, **kwargs) + + +def dinov2_vitg14(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-g/14 model (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model( + arch_name="vit_giant2", + ffn_layer="swiglufused", + weights=weights, + pretrained=pretrained, + **kwargs, + ) + + +def dinov2_vits14_reg(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-S/14 model with registers (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model( + arch_name="vit_small", + pretrained=pretrained, + weights=weights, + num_register_tokens=4, + interpolate_antialias=True, + interpolate_offset=0.0, + **kwargs, + ) + + +def dinov2_vitb14_reg(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-B/14 model with registers (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model( + arch_name="vit_base", + pretrained=pretrained, + weights=weights, + num_register_tokens=4, + interpolate_antialias=True, + interpolate_offset=0.0, + **kwargs, + ) + + +def dinov2_vitl14_reg(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-L/14 model with registers (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model( + arch_name="vit_large", + pretrained=pretrained, + weights=weights, + num_register_tokens=4, + interpolate_antialias=True, + interpolate_offset=0.0, + **kwargs, + ) + + +def dinov2_vitg14_reg(*, pretrained: bool = True, weights: Union[Weights, str] = Weights.LVD142M, **kwargs): + """ + DINOv2 ViT-g/14 model with registers (optionally) pretrained on the LVD-142M dataset. + """ + return _make_dinov2_model( + arch_name="vit_giant2", + ffn_layer="swiglufused", + weights=weights, + pretrained=pretrained, + num_register_tokens=4, + interpolate_antialias=True, + interpolate_offset=0.0, + **kwargs, + ) diff --git a/moge/model/dinov2/hub/utils.py b/moge/model/dinov2/hub/utils.py new file mode 100644 index 0000000..9c66414 --- /dev/null +++ b/moge/model/dinov2/hub/utils.py @@ -0,0 +1,39 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import itertools +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +_DINOV2_BASE_URL = "https://dl.fbaipublicfiles.com/dinov2" + + +def _make_dinov2_model_name(arch_name: str, patch_size: int, num_register_tokens: int = 0) -> str: + compact_arch_name = arch_name.replace("_", "")[:4] + registers_suffix = f"_reg{num_register_tokens}" if num_register_tokens else "" + return f"dinov2_{compact_arch_name}{patch_size}{registers_suffix}" + + +class CenterPadding(nn.Module): + def __init__(self, multiple): + super().__init__() + self.multiple = multiple + + def _get_pad(self, size): + new_size = math.ceil(size / self.multiple) * self.multiple + pad_size = new_size - size + pad_size_left = pad_size // 2 + pad_size_right = pad_size - pad_size_left + return pad_size_left, pad_size_right + + @torch.inference_mode() + def forward(self, x): + pads = list(itertools.chain.from_iterable(self._get_pad(m) for m in x.shape[:1:-1])) + output = F.pad(x, pads) + return output diff --git a/moge/model/dinov2/layers/__init__.py b/moge/model/dinov2/layers/__init__.py new file mode 100644 index 0000000..05a0b61 --- /dev/null +++ b/moge/model/dinov2/layers/__init__.py @@ -0,0 +1,11 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +from .dino_head import DINOHead +from .mlp import Mlp +from .patch_embed import PatchEmbed +from .swiglu_ffn import SwiGLUFFN, SwiGLUFFNFused +from .block import NestedTensorBlock +from .attention import MemEffAttention diff --git a/moge/model/dinov2/layers/attention.py b/moge/model/dinov2/layers/attention.py new file mode 100644 index 0000000..c9f79d4 --- /dev/null +++ b/moge/model/dinov2/layers/attention.py @@ -0,0 +1,100 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/models/vision_transformer.py + +import logging +import os +import warnings + +import torch.nn.functional as F +from torch import Tensor +from torch import nn + + +logger = logging.getLogger("dinov2") + + +XFORMERS_ENABLED = os.environ.get("XFORMERS_DISABLED") is None +try: + if XFORMERS_ENABLED: + from xformers.ops import memory_efficient_attention, unbind + + XFORMERS_AVAILABLE = True + # warnings.warn("xFormers is available (Attention)") + else: + # warnings.warn("xFormers is disabled (Attention)") + raise ImportError +except ImportError: + XFORMERS_AVAILABLE = False + # warnings.warn("xFormers is not available (Attention)") + + +class Attention(nn.Module): + def __init__( + self, + dim: int, + num_heads: int = 8, + qkv_bias: bool = False, + proj_bias: bool = True, + attn_drop: float = 0.0, + proj_drop: float = 0.0, + ) -> None: + super().__init__() + self.num_heads = num_heads + head_dim = dim // num_heads + self.scale = head_dim**-0.5 + + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(dim, dim, bias=proj_bias) + self.proj_drop = nn.Dropout(proj_drop) + + # # Deprecated implementation, extremely slow + # def forward(self, x: Tensor, attn_bias=None) -> Tensor: + # B, N, C = x.shape + # qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) + # q, k, v = qkv[0] * self.scale, qkv[1], qkv[2] + # attn = q @ k.transpose(-2, -1) + # attn = attn.softmax(dim=-1) + # attn = self.attn_drop(attn) + # x = (attn @ v).transpose(1, 2).reshape(B, N, C) + # x = self.proj(x) + # x = self.proj_drop(x) + # return x + + def forward(self, x: Tensor, attn_bias=None) -> Tensor: + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) # (3, B, H, N, C // H) + + q, k, v = qkv.unbind(0) # (B, H, N, C // H) + + x = F.scaled_dot_product_attention(q, k, v, attn_bias) + x = x.permute(0, 2, 1, 3).reshape(B, N, C) + + x = self.proj(x) + x = self.proj_drop(x) + return x + +class MemEffAttention(Attention): + def forward(self, x: Tensor, attn_bias=None) -> Tensor: + if not XFORMERS_AVAILABLE: + if attn_bias is not None: + raise AssertionError("xFormers is required for using nested tensors") + return super().forward(x) + + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) + + q, k, v = unbind(qkv, 2) + + x = memory_efficient_attention(q, k, v, attn_bias=attn_bias) + x = x.reshape([B, N, C]) + + x = self.proj(x) + x = self.proj_drop(x) + return x diff --git a/moge/model/dinov2/layers/block.py b/moge/model/dinov2/layers/block.py new file mode 100644 index 0000000..fd5b8a7 --- /dev/null +++ b/moge/model/dinov2/layers/block.py @@ -0,0 +1,259 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/patch_embed.py + +import logging +import os +from typing import Callable, List, Any, Tuple, Dict +import warnings + +import torch +from torch import nn, Tensor + +from .attention import Attention, MemEffAttention +from .drop_path import DropPath +from .layer_scale import LayerScale +from .mlp import Mlp + + +logger = logging.getLogger("dinov2") + + +XFORMERS_ENABLED = os.environ.get("XFORMERS_DISABLED") is None +try: + if XFORMERS_ENABLED: + from xformers.ops import fmha, scaled_index_add, index_select_cat + + XFORMERS_AVAILABLE = True + # warnings.warn("xFormers is available (Block)") + else: + # warnings.warn("xFormers is disabled (Block)") + raise ImportError +except ImportError: + XFORMERS_AVAILABLE = False + # warnings.warn("xFormers is not available (Block)") + + +class Block(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + mlp_ratio: float = 4.0, + qkv_bias: bool = False, + proj_bias: bool = True, + ffn_bias: bool = True, + drop: float = 0.0, + attn_drop: float = 0.0, + init_values=None, + drop_path: float = 0.0, + act_layer: Callable[..., nn.Module] = nn.GELU, + norm_layer: Callable[..., nn.Module] = nn.LayerNorm, + attn_class: Callable[..., nn.Module] = Attention, + ffn_layer: Callable[..., nn.Module] = Mlp, + ) -> None: + super().__init__() + # print(f"biases: qkv: {qkv_bias}, proj: {proj_bias}, ffn: {ffn_bias}") + self.norm1 = norm_layer(dim) + self.attn = attn_class( + dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + attn_drop=attn_drop, + proj_drop=drop, + ) + self.ls1 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity() + self.drop_path1 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + + self.norm2 = norm_layer(dim) + mlp_hidden_dim = int(dim * mlp_ratio) + self.mlp = ffn_layer( + in_features=dim, + hidden_features=mlp_hidden_dim, + act_layer=act_layer, + drop=drop, + bias=ffn_bias, + ) + self.ls2 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity() + self.drop_path2 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + + self.sample_drop_ratio = drop_path + + def forward(self, x: Tensor) -> Tensor: + def attn_residual_func(x: Tensor) -> Tensor: + return self.ls1(self.attn(self.norm1(x))) + + def ffn_residual_func(x: Tensor) -> Tensor: + return self.ls2(self.mlp(self.norm2(x))) + + if self.training and self.sample_drop_ratio > 0.1: + # the overhead is compensated only for a drop path rate larger than 0.1 + x = drop_add_residual_stochastic_depth( + x, + residual_func=attn_residual_func, + sample_drop_ratio=self.sample_drop_ratio, + ) + x = drop_add_residual_stochastic_depth( + x, + residual_func=ffn_residual_func, + sample_drop_ratio=self.sample_drop_ratio, + ) + elif self.training and self.sample_drop_ratio > 0.0: + x = x + self.drop_path1(attn_residual_func(x)) + x = x + self.drop_path1(ffn_residual_func(x)) # FIXME: drop_path2 + else: + x = x + attn_residual_func(x) + x = x + ffn_residual_func(x) + return x + + +def drop_add_residual_stochastic_depth( + x: Tensor, + residual_func: Callable[[Tensor], Tensor], + sample_drop_ratio: float = 0.0, +) -> Tensor: + # 1) extract subset using permutation + b, n, d = x.shape + sample_subset_size = max(int(b * (1 - sample_drop_ratio)), 1) + brange = (torch.randperm(b, device=x.device))[:sample_subset_size] + x_subset = x[brange] + + # 2) apply residual_func to get residual + residual = residual_func(x_subset) + + x_flat = x.flatten(1) + residual = residual.flatten(1) + + residual_scale_factor = b / sample_subset_size + + # 3) add the residual + x_plus_residual = torch.index_add(x_flat, 0, brange, residual.to(dtype=x.dtype), alpha=residual_scale_factor) + return x_plus_residual.view_as(x) + + +def get_branges_scales(x, sample_drop_ratio=0.0): + b, n, d = x.shape + sample_subset_size = max(int(b * (1 - sample_drop_ratio)), 1) + brange = (torch.randperm(b, device=x.device))[:sample_subset_size] + residual_scale_factor = b / sample_subset_size + return brange, residual_scale_factor + + +def add_residual(x, brange, residual, residual_scale_factor, scaling_vector=None): + if scaling_vector is None: + x_flat = x.flatten(1) + residual = residual.flatten(1) + x_plus_residual = torch.index_add(x_flat, 0, brange, residual.to(dtype=x.dtype), alpha=residual_scale_factor) + else: + x_plus_residual = scaled_index_add( + x, brange, residual.to(dtype=x.dtype), scaling=scaling_vector, alpha=residual_scale_factor + ) + return x_plus_residual + + +attn_bias_cache: Dict[Tuple, Any] = {} + + +def get_attn_bias_and_cat(x_list, branges=None): + """ + this will perform the index select, cat the tensors, and provide the attn_bias from cache + """ + batch_sizes = [b.shape[0] for b in branges] if branges is not None else [x.shape[0] for x in x_list] + all_shapes = tuple((b, x.shape[1]) for b, x in zip(batch_sizes, x_list)) + if all_shapes not in attn_bias_cache.keys(): + seqlens = [] + for b, x in zip(batch_sizes, x_list): + for _ in range(b): + seqlens.append(x.shape[1]) + attn_bias = fmha.BlockDiagonalMask.from_seqlens(seqlens) + attn_bias._batch_sizes = batch_sizes + attn_bias_cache[all_shapes] = attn_bias + + if branges is not None: + cat_tensors = index_select_cat([x.flatten(1) for x in x_list], branges).view(1, -1, x_list[0].shape[-1]) + else: + tensors_bs1 = tuple(x.reshape([1, -1, *x.shape[2:]]) for x in x_list) + cat_tensors = torch.cat(tensors_bs1, dim=1) + + return attn_bias_cache[all_shapes], cat_tensors + + +def drop_add_residual_stochastic_depth_list( + x_list: List[Tensor], + residual_func: Callable[[Tensor, Any], Tensor], + sample_drop_ratio: float = 0.0, + scaling_vector=None, +) -> Tensor: + # 1) generate random set of indices for dropping samples in the batch + branges_scales = [get_branges_scales(x, sample_drop_ratio=sample_drop_ratio) for x in x_list] + branges = [s[0] for s in branges_scales] + residual_scale_factors = [s[1] for s in branges_scales] + + # 2) get attention bias and index+concat the tensors + attn_bias, x_cat = get_attn_bias_and_cat(x_list, branges) + + # 3) apply residual_func to get residual, and split the result + residual_list = attn_bias.split(residual_func(x_cat, attn_bias=attn_bias)) # type: ignore + + outputs = [] + for x, brange, residual, residual_scale_factor in zip(x_list, branges, residual_list, residual_scale_factors): + outputs.append(add_residual(x, brange, residual, residual_scale_factor, scaling_vector).view_as(x)) + return outputs + + +class NestedTensorBlock(Block): + def forward_nested(self, x_list: List[Tensor]) -> List[Tensor]: + """ + x_list contains a list of tensors to nest together and run + """ + assert isinstance(self.attn, MemEffAttention) + + if self.training and self.sample_drop_ratio > 0.0: + + def attn_residual_func(x: Tensor, attn_bias=None) -> Tensor: + return self.attn(self.norm1(x), attn_bias=attn_bias) + + def ffn_residual_func(x: Tensor, attn_bias=None) -> Tensor: + return self.mlp(self.norm2(x)) + + x_list = drop_add_residual_stochastic_depth_list( + x_list, + residual_func=attn_residual_func, + sample_drop_ratio=self.sample_drop_ratio, + scaling_vector=self.ls1.gamma if isinstance(self.ls1, LayerScale) else None, + ) + x_list = drop_add_residual_stochastic_depth_list( + x_list, + residual_func=ffn_residual_func, + sample_drop_ratio=self.sample_drop_ratio, + scaling_vector=self.ls2.gamma if isinstance(self.ls1, LayerScale) else None, + ) + return x_list + else: + + def attn_residual_func(x: Tensor, attn_bias=None) -> Tensor: + return self.ls1(self.attn(self.norm1(x), attn_bias=attn_bias)) + + def ffn_residual_func(x: Tensor, attn_bias=None) -> Tensor: + return self.ls2(self.mlp(self.norm2(x))) + + attn_bias, x = get_attn_bias_and_cat(x_list) + x = x + attn_residual_func(x, attn_bias=attn_bias) + x = x + ffn_residual_func(x) + return attn_bias.split(x) + + def forward(self, x_or_x_list): + if isinstance(x_or_x_list, Tensor): + return super().forward(x_or_x_list) + elif isinstance(x_or_x_list, list): + if not XFORMERS_AVAILABLE: + raise AssertionError("xFormers is required for using nested tensors") + return self.forward_nested(x_or_x_list) + else: + raise AssertionError diff --git a/moge/model/dinov2/layers/dino_head.py b/moge/model/dinov2/layers/dino_head.py new file mode 100644 index 0000000..0ace8ff --- /dev/null +++ b/moge/model/dinov2/layers/dino_head.py @@ -0,0 +1,58 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import torch +import torch.nn as nn +from torch.nn.init import trunc_normal_ +from torch.nn.utils import weight_norm + + +class DINOHead(nn.Module): + def __init__( + self, + in_dim, + out_dim, + use_bn=False, + nlayers=3, + hidden_dim=2048, + bottleneck_dim=256, + mlp_bias=True, + ): + super().__init__() + nlayers = max(nlayers, 1) + self.mlp = _build_mlp(nlayers, in_dim, bottleneck_dim, hidden_dim=hidden_dim, use_bn=use_bn, bias=mlp_bias) + self.apply(self._init_weights) + self.last_layer = weight_norm(nn.Linear(bottleneck_dim, out_dim, bias=False)) + self.last_layer.weight_g.data.fill_(1) + + def _init_weights(self, m): + if isinstance(m, nn.Linear): + trunc_normal_(m.weight, std=0.02) + if isinstance(m, nn.Linear) and m.bias is not None: + nn.init.constant_(m.bias, 0) + + def forward(self, x): + x = self.mlp(x) + eps = 1e-6 if x.dtype == torch.float16 else 1e-12 + x = nn.functional.normalize(x, dim=-1, p=2, eps=eps) + x = self.last_layer(x) + return x + + +def _build_mlp(nlayers, in_dim, bottleneck_dim, hidden_dim=None, use_bn=False, bias=True): + if nlayers == 1: + return nn.Linear(in_dim, bottleneck_dim, bias=bias) + else: + layers = [nn.Linear(in_dim, hidden_dim, bias=bias)] + if use_bn: + layers.append(nn.BatchNorm1d(hidden_dim)) + layers.append(nn.GELU()) + for _ in range(nlayers - 2): + layers.append(nn.Linear(hidden_dim, hidden_dim, bias=bias)) + if use_bn: + layers.append(nn.BatchNorm1d(hidden_dim)) + layers.append(nn.GELU()) + layers.append(nn.Linear(hidden_dim, bottleneck_dim, bias=bias)) + return nn.Sequential(*layers) diff --git a/moge/model/dinov2/layers/drop_path.py b/moge/model/dinov2/layers/drop_path.py new file mode 100644 index 0000000..1d640e0 --- /dev/null +++ b/moge/model/dinov2/layers/drop_path.py @@ -0,0 +1,34 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/drop.py + + +from torch import nn + + +def drop_path(x, drop_prob: float = 0.0, training: bool = False): + if drop_prob == 0.0 or not training: + return x + keep_prob = 1 - drop_prob + shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets + random_tensor = x.new_empty(shape).bernoulli_(keep_prob) + if keep_prob > 0.0: + random_tensor.div_(keep_prob) + output = x * random_tensor + return output + + +class DropPath(nn.Module): + """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).""" + + def __init__(self, drop_prob=None): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + + def forward(self, x): + return drop_path(x, self.drop_prob, self.training) diff --git a/moge/model/dinov2/layers/layer_scale.py b/moge/model/dinov2/layers/layer_scale.py new file mode 100644 index 0000000..51df0d7 --- /dev/null +++ b/moge/model/dinov2/layers/layer_scale.py @@ -0,0 +1,27 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# Modified from: https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py#L103-L110 + +from typing import Union + +import torch +from torch import Tensor +from torch import nn + + +class LayerScale(nn.Module): + def __init__( + self, + dim: int, + init_values: Union[float, Tensor] = 1e-5, + inplace: bool = False, + ) -> None: + super().__init__() + self.inplace = inplace + self.gamma = nn.Parameter(init_values * torch.ones(dim)) + + def forward(self, x: Tensor) -> Tensor: + return x.mul_(self.gamma) if self.inplace else x * self.gamma diff --git a/moge/model/dinov2/layers/mlp.py b/moge/model/dinov2/layers/mlp.py new file mode 100644 index 0000000..bbf9432 --- /dev/null +++ b/moge/model/dinov2/layers/mlp.py @@ -0,0 +1,40 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/mlp.py + + +from typing import Callable, Optional + +from torch import Tensor, nn + + +class Mlp(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: Optional[int] = None, + out_features: Optional[int] = None, + act_layer: Callable[..., nn.Module] = nn.GELU, + drop: float = 0.0, + bias: bool = True, + ) -> None: + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.fc1 = nn.Linear(in_features, hidden_features, bias=bias) + self.act = act_layer() + self.fc2 = nn.Linear(hidden_features, out_features, bias=bias) + self.drop = nn.Dropout(drop) + + def forward(self, x: Tensor) -> Tensor: + x = self.fc1(x) + x = self.act(x) + x = self.drop(x) + x = self.fc2(x) + x = self.drop(x) + return x diff --git a/moge/model/dinov2/layers/patch_embed.py b/moge/model/dinov2/layers/patch_embed.py new file mode 100644 index 0000000..8b7c080 --- /dev/null +++ b/moge/model/dinov2/layers/patch_embed.py @@ -0,0 +1,88 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/patch_embed.py + +from typing import Callable, Optional, Tuple, Union + +from torch import Tensor +import torch.nn as nn + + +def make_2tuple(x): + if isinstance(x, tuple): + assert len(x) == 2 + return x + + assert isinstance(x, int) + return (x, x) + + +class PatchEmbed(nn.Module): + """ + 2D image to patch embedding: (B,C,H,W) -> (B,N,D) + + Args: + img_size: Image size. + patch_size: Patch token size. + in_chans: Number of input image channels. + embed_dim: Number of linear projection output channels. + norm_layer: Normalization layer. + """ + + def __init__( + self, + img_size: Union[int, Tuple[int, int]] = 224, + patch_size: Union[int, Tuple[int, int]] = 16, + in_chans: int = 3, + embed_dim: int = 768, + norm_layer: Optional[Callable] = None, + flatten_embedding: bool = True, + ) -> None: + super().__init__() + + image_HW = make_2tuple(img_size) + patch_HW = make_2tuple(patch_size) + patch_grid_size = ( + image_HW[0] // patch_HW[0], + image_HW[1] // patch_HW[1], + ) + + self.img_size = image_HW + self.patch_size = patch_HW + self.patches_resolution = patch_grid_size + self.num_patches = patch_grid_size[0] * patch_grid_size[1] + + self.in_chans = in_chans + self.embed_dim = embed_dim + + self.flatten_embedding = flatten_embedding + + self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_HW, stride=patch_HW) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x: Tensor) -> Tensor: + _, _, H, W = x.shape + patch_H, patch_W = self.patch_size + + assert H % patch_H == 0, f"Input image height {H} is not a multiple of patch height {patch_H}" + assert W % patch_W == 0, f"Input image width {W} is not a multiple of patch width: {patch_W}" + + x = self.proj(x) # B C H W + H, W = x.size(2), x.size(3) + x = x.flatten(2).transpose(1, 2) # B HW C + x = self.norm(x) + if not self.flatten_embedding: + x = x.reshape(-1, H, W, self.embed_dim) # B H W C + return x + + def flops(self) -> float: + Ho, Wo = self.patches_resolution + flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1]) + if self.norm is not None: + flops += Ho * Wo * self.embed_dim + return flops diff --git a/moge/model/dinov2/layers/swiglu_ffn.py b/moge/model/dinov2/layers/swiglu_ffn.py new file mode 100644 index 0000000..5ce2115 --- /dev/null +++ b/moge/model/dinov2/layers/swiglu_ffn.py @@ -0,0 +1,72 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import os +from typing import Callable, Optional +import warnings + +from torch import Tensor, nn +import torch.nn.functional as F + + +class SwiGLUFFN(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: Optional[int] = None, + out_features: Optional[int] = None, + act_layer: Callable[..., nn.Module] = None, + drop: float = 0.0, + bias: bool = True, + ) -> None: + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.w12 = nn.Linear(in_features, 2 * hidden_features, bias=bias) + self.w3 = nn.Linear(hidden_features, out_features, bias=bias) + + def forward(self, x: Tensor) -> Tensor: + x12 = self.w12(x) + x1, x2 = x12.chunk(2, dim=-1) + hidden = F.silu(x1) * x2 + return self.w3(hidden) + + +XFORMERS_ENABLED = os.environ.get("XFORMERS_DISABLED") is None +try: + if XFORMERS_ENABLED: + from xformers.ops import SwiGLU + + XFORMERS_AVAILABLE = True + # warnings.warn("xFormers is available (SwiGLU)") + else: + # warnings.warn("xFormers is disabled (SwiGLU)") + raise ImportError +except ImportError: + SwiGLU = SwiGLUFFN + XFORMERS_AVAILABLE = False + + # warnings.warn("xFormers is not available (SwiGLU)") + + +class SwiGLUFFNFused(SwiGLU): + def __init__( + self, + in_features: int, + hidden_features: Optional[int] = None, + out_features: Optional[int] = None, + act_layer: Callable[..., nn.Module] = None, + drop: float = 0.0, + bias: bool = True, + ) -> None: + out_features = out_features or in_features + hidden_features = hidden_features or in_features + hidden_features = (int(hidden_features * 2 / 3) + 7) // 8 * 8 + super().__init__( + in_features=in_features, + hidden_features=hidden_features, + out_features=out_features, + bias=bias, + ) diff --git a/moge/model/dinov2/models/__init__.py b/moge/model/dinov2/models/__init__.py new file mode 100644 index 0000000..3fdff20 --- /dev/null +++ b/moge/model/dinov2/models/__init__.py @@ -0,0 +1,43 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import logging + +from . import vision_transformer as vits + + +logger = logging.getLogger("dinov2") + + +def build_model(args, only_teacher=False, img_size=224): + args.arch = args.arch.removesuffix("_memeff") + if "vit" in args.arch: + vit_kwargs = dict( + img_size=img_size, + patch_size=args.patch_size, + init_values=args.layerscale, + ffn_layer=args.ffn_layer, + block_chunks=args.block_chunks, + qkv_bias=args.qkv_bias, + proj_bias=args.proj_bias, + ffn_bias=args.ffn_bias, + num_register_tokens=args.num_register_tokens, + interpolate_offset=args.interpolate_offset, + interpolate_antialias=args.interpolate_antialias, + ) + teacher = vits.__dict__[args.arch](**vit_kwargs) + if only_teacher: + return teacher, teacher.embed_dim + student = vits.__dict__[args.arch]( + **vit_kwargs, + drop_path_rate=args.drop_path_rate, + drop_path_uniform=args.drop_path_uniform, + ) + embed_dim = student.embed_dim + return student, teacher, embed_dim + + +def build_model_from_cfg(cfg, only_teacher=False): + return build_model(cfg.student, only_teacher=only_teacher, img_size=cfg.crops.global_crops_size) diff --git a/moge/model/dinov2/models/vision_transformer.py b/moge/model/dinov2/models/vision_transformer.py new file mode 100644 index 0000000..f0bed9d --- /dev/null +++ b/moge/model/dinov2/models/vision_transformer.py @@ -0,0 +1,407 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +# References: +# https://github.com/facebookresearch/dino/blob/main/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/models/vision_transformer.py + +from functools import partial +import math +import logging +from typing import Sequence, Tuple, Union, Callable, Optional, List + +import torch +import torch.nn as nn +import torch.utils.checkpoint +from torch.nn.init import trunc_normal_ + +from ..layers import Mlp, PatchEmbed, SwiGLUFFNFused, MemEffAttention, NestedTensorBlock as Block + + +logger = logging.getLogger("dinov2") + + +def named_apply(fn: Callable, module: nn.Module, name="", depth_first=True, include_root=False) -> nn.Module: + if not depth_first and include_root: + fn(module=module, name=name) + for child_name, child_module in module.named_children(): + child_name = ".".join((name, child_name)) if name else child_name + named_apply(fn=fn, module=child_module, name=child_name, depth_first=depth_first, include_root=True) + if depth_first and include_root: + fn(module=module, name=name) + return module + + +class BlockChunk(nn.ModuleList): + def forward(self, x): + for b in self: + x = b(x) + return x + + +class DinoVisionTransformer(nn.Module): + def __init__( + self, + img_size=224, + patch_size=16, + in_chans=3, + embed_dim=768, + depth=12, + num_heads=12, + mlp_ratio=4.0, + qkv_bias=True, + ffn_bias=True, + proj_bias=True, + drop_path_rate=0.0, + drop_path_uniform=False, + init_values=None, # for layerscale: None or 0 => no layerscale + embed_layer=PatchEmbed, + act_layer=nn.GELU, + block_fn=Block, + ffn_layer="mlp", + block_chunks=1, + num_register_tokens=0, + interpolate_antialias=False, + interpolate_offset=0.1, + ): + """ + Args: + img_size (int, tuple): input image size + patch_size (int, tuple): patch size + in_chans (int): number of input channels + embed_dim (int): embedding dimension + depth (int): depth of transformer + num_heads (int): number of attention heads + mlp_ratio (int): ratio of mlp hidden dim to embedding dim + qkv_bias (bool): enable bias for qkv if True + proj_bias (bool): enable bias for proj in attn if True + ffn_bias (bool): enable bias for ffn if True + drop_path_rate (float): stochastic depth rate + drop_path_uniform (bool): apply uniform drop rate across blocks + weight_init (str): weight init scheme + init_values (float): layer-scale init values + embed_layer (nn.Module): patch embedding layer + act_layer (nn.Module): MLP activation layer + block_fn (nn.Module): transformer block class + ffn_layer (str): "mlp", "swiglu", "swiglufused" or "identity" + block_chunks: (int) split block sequence into block_chunks units for FSDP wrap + num_register_tokens: (int) number of extra cls tokens (so-called "registers") + interpolate_antialias: (str) flag to apply anti-aliasing when interpolating positional embeddings + interpolate_offset: (float) work-around offset to apply when interpolating positional embeddings + """ + super().__init__() + norm_layer = partial(nn.LayerNorm, eps=1e-6) + + self.num_features = self.embed_dim = embed_dim # num_features for consistency with other models + self.num_tokens = 1 + self.n_blocks = depth + self.num_heads = num_heads + self.patch_size = patch_size + self.num_register_tokens = num_register_tokens + self.interpolate_antialias = interpolate_antialias + self.interpolate_offset = interpolate_offset + + self.patch_embed = embed_layer(img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim) + num_patches = self.patch_embed.num_patches + + self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) + self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + self.num_tokens, embed_dim)) + assert num_register_tokens >= 0 + self.register_tokens = ( + nn.Parameter(torch.zeros(1, num_register_tokens, embed_dim)) if num_register_tokens else None + ) + + if drop_path_uniform is True: + dpr = [drop_path_rate] * depth + else: + dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule + + if ffn_layer == "mlp": + logger.info("using MLP layer as FFN") + ffn_layer = Mlp + elif ffn_layer == "swiglufused" or ffn_layer == "swiglu": + logger.info("using SwiGLU layer as FFN") + ffn_layer = SwiGLUFFNFused + elif ffn_layer == "identity": + logger.info("using Identity layer as FFN") + + def f(*args, **kwargs): + return nn.Identity() + + ffn_layer = f + else: + raise NotImplementedError + + blocks_list = [ + block_fn( + dim=embed_dim, + num_heads=num_heads, + mlp_ratio=mlp_ratio, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + ffn_bias=ffn_bias, + drop_path=dpr[i], + norm_layer=norm_layer, + act_layer=act_layer, + ffn_layer=ffn_layer, + init_values=init_values, + ) + for i in range(depth) + ] + if block_chunks > 0: + self.chunked_blocks = True + chunked_blocks = [] + chunksize = depth // block_chunks + for i in range(0, depth, chunksize): + # this is to keep the block index consistent if we chunk the block list + chunked_blocks.append([nn.Identity()] * i + blocks_list[i : i + chunksize]) + self.blocks = nn.ModuleList([BlockChunk(p) for p in chunked_blocks]) + else: + self.chunked_blocks = False + self.blocks = nn.ModuleList(blocks_list) + + self.norm = norm_layer(embed_dim) + self.head = nn.Identity() + + self.mask_token = nn.Parameter(torch.zeros(1, embed_dim)) + + self.init_weights() + + @property + def onnx_compatible_mode(self): + return getattr(self, "_onnx_compatible_mode", False) + + @onnx_compatible_mode.setter + def onnx_compatible_mode(self, value: bool): + self._onnx_compatible_mode = value + + def init_weights(self): + trunc_normal_(self.pos_embed, std=0.02) + nn.init.normal_(self.cls_token, std=1e-6) + if self.register_tokens is not None: + nn.init.normal_(self.register_tokens, std=1e-6) + named_apply(init_weights_vit_timm, self) + + def interpolate_pos_encoding(self, x, h, w): + previous_dtype = x.dtype + npatch = x.shape[1] - 1 + batch_size = x.shape[0] + N = self.pos_embed.shape[1] - 1 + if not self.onnx_compatible_mode and npatch == N and w == h: + return self.pos_embed + pos_embed = self.pos_embed.float() + class_pos_embed = pos_embed[:, 0, :] + patch_pos_embed = pos_embed[:, 1:, :] + dim = x.shape[-1] + h0, w0 = h // self.patch_size, w // self.patch_size + M = int(math.sqrt(N)) # Recover the number of patches in each dimension + assert N == M * M + kwargs = {} + if not self.onnx_compatible_mode and self.interpolate_offset > 0: + # Historical kludge: add a small number to avoid floating point error in the interpolation, see https://github.com/facebookresearch/dino/issues/8 + # Note: still needed for backward-compatibility, the underlying operators are using both output size and scale factors + sx = float(w0 + self.interpolate_offset) / M + sy = float(h0 + self.interpolate_offset) / M + kwargs["scale_factor"] = (sy, sx) + else: + # Simply specify an output size instead of a scale factor + kwargs["size"] = (h0, w0) + + patch_pos_embed = nn.functional.interpolate( + patch_pos_embed.reshape(1, M, M, dim).permute(0, 3, 1, 2), + mode="bicubic", + antialias=self.interpolate_antialias, + **kwargs, + ) + + assert (h0, w0) == patch_pos_embed.shape[-2:] + patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).flatten(1, 2) + return torch.cat((class_pos_embed[:, None, :].expand(patch_pos_embed.shape[0], -1, -1), patch_pos_embed), dim=1).to(previous_dtype) + + def prepare_tokens_with_masks(self, x, masks=None): + B, nc, h, w = x.shape + x = self.patch_embed(x) + + if masks is not None: + x = torch.where(masks.unsqueeze(-1), self.mask_token.to(x.dtype).unsqueeze(0), x) + + x = torch.cat((self.cls_token.expand(x.shape[0], -1, -1), x), dim=1) + x = x + self.interpolate_pos_encoding(x, h, w) + + if self.register_tokens is not None: + x = torch.cat( + ( + x[:, :1], + self.register_tokens.expand(x.shape[0], -1, -1), + x[:, 1:], + ), + dim=1, + ) + + return x + + def forward_features_list(self, x_list, masks_list): + x = [self.prepare_tokens_with_masks(x, masks) for x, masks, ar in zip(x_list, masks_list)] + for blk in self.blocks: + x = blk(x) + + all_x = x + output = [] + for x, masks in zip(all_x, masks_list): + x_norm = self.norm(x) + output.append( + { + "x_norm_clstoken": x_norm[:, 0], + "x_norm_regtokens": x_norm[:, 1 : self.num_register_tokens + 1], + "x_norm_patchtokens": x_norm[:, self.num_register_tokens + 1 :], + "x_prenorm": x, + "masks": masks, + } + ) + return output + + def forward_features(self, x, masks=None): + if isinstance(x, list): + return self.forward_features_list(x, masks) + + x = self.prepare_tokens_with_masks(x, masks) + + for blk in self.blocks: + x = blk(x) + + x_norm = self.norm(x) + return { + "x_norm_clstoken": x_norm[:, 0], + "x_norm_regtokens": x_norm[:, 1 : self.num_register_tokens + 1], + "x_norm_patchtokens": x_norm[:, self.num_register_tokens + 1 :], + "x_prenorm": x, + "masks": masks, + } + + def _get_intermediate_layers_not_chunked(self, x, n=1): + x = self.prepare_tokens_with_masks(x) + # If n is an int, take the n last blocks. If it's a list, take them + output, total_block_len = [], len(self.blocks) + blocks_to_take = range(total_block_len - n, total_block_len) if isinstance(n, int) else n + for i, blk in enumerate(self.blocks): + x = blk(x) + if i in blocks_to_take: + output.append(x) + assert len(output) == len(blocks_to_take), f"only {len(output)} / {len(blocks_to_take)} blocks found" + return output + + def _get_intermediate_layers_chunked(self, x, n=1): + x = self.prepare_tokens_with_masks(x) + output, i, total_block_len = [], 0, len(self.blocks[-1]) + # If n is an int, take the n last blocks. If it's a list, take them + blocks_to_take = range(total_block_len - n, total_block_len) if isinstance(n, int) else n + for block_chunk in self.blocks: + for blk in block_chunk[i:]: # Passing the nn.Identity() + x = blk(x) + if i in blocks_to_take: + output.append(x) + i += 1 + assert len(output) == len(blocks_to_take), f"only {len(output)} / {len(blocks_to_take)} blocks found" + return output + + def get_intermediate_layers( + self, + x: torch.Tensor, + n: Union[int, Sequence] = 1, # Layers or n last layers to take + reshape: bool = False, + return_class_token: bool = False, + norm=True, + ) -> Tuple[Union[torch.Tensor, Tuple[torch.Tensor]]]: + if self.chunked_blocks: + outputs = self._get_intermediate_layers_chunked(x, n) + else: + outputs = self._get_intermediate_layers_not_chunked(x, n) + if norm: + outputs = [self.norm(out) for out in outputs] + class_tokens = [out[:, 0] for out in outputs] + outputs = [out[:, 1 + self.num_register_tokens :] for out in outputs] + if reshape: + B, _, w, h = x.shape + outputs = [ + out.reshape(B, w // self.patch_size, h // self.patch_size, -1).permute(0, 3, 1, 2).contiguous() + for out in outputs + ] + if return_class_token: + return tuple(zip(outputs, class_tokens)) + return tuple(outputs) + + def forward(self, *args, is_training=False, **kwargs): + ret = self.forward_features(*args, **kwargs) + if is_training: + return ret + else: + return self.head(ret["x_norm_clstoken"]) + + +def init_weights_vit_timm(module: nn.Module, name: str = ""): + """ViT weight initialization, original timm impl (for reproducibility)""" + if isinstance(module, nn.Linear): + trunc_normal_(module.weight, std=0.02) + if module.bias is not None: + nn.init.zeros_(module.bias) + + +def vit_small(patch_size=16, num_register_tokens=0, **kwargs): + model = DinoVisionTransformer( + patch_size=patch_size, + embed_dim=384, + depth=12, + num_heads=6, + mlp_ratio=4, + block_fn=partial(Block, attn_class=MemEffAttention), + num_register_tokens=num_register_tokens, + **kwargs, + ) + return model + + +def vit_base(patch_size=16, num_register_tokens=0, **kwargs): + model = DinoVisionTransformer( + patch_size=patch_size, + embed_dim=768, + depth=12, + num_heads=12, + mlp_ratio=4, + block_fn=partial(Block, attn_class=MemEffAttention), + num_register_tokens=num_register_tokens, + **kwargs, + ) + return model + + +def vit_large(patch_size=16, num_register_tokens=0, **kwargs): + model = DinoVisionTransformer( + patch_size=patch_size, + embed_dim=1024, + depth=24, + num_heads=16, + mlp_ratio=4, + block_fn=partial(Block, attn_class=MemEffAttention), + num_register_tokens=num_register_tokens, + **kwargs, + ) + return model + + +def vit_giant2(patch_size=16, num_register_tokens=0, **kwargs): + """ + Close to ViT-giant, with embed-dim 1536 and 24 heads => embed-dim per head 64 + """ + model = DinoVisionTransformer( + patch_size=patch_size, + embed_dim=1536, + depth=40, + num_heads=24, + mlp_ratio=4, + block_fn=partial(Block, attn_class=MemEffAttention), + num_register_tokens=num_register_tokens, + **kwargs, + ) + return model diff --git a/moge/model/dinov2/utils/__init__.py b/moge/model/dinov2/utils/__init__.py new file mode 100644 index 0000000..b88da6b --- /dev/null +++ b/moge/model/dinov2/utils/__init__.py @@ -0,0 +1,4 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. diff --git a/moge/model/dinov2/utils/cluster.py b/moge/model/dinov2/utils/cluster.py new file mode 100644 index 0000000..3df87dc --- /dev/null +++ b/moge/model/dinov2/utils/cluster.py @@ -0,0 +1,95 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +from enum import Enum +import os +from pathlib import Path +from typing import Any, Dict, Optional + + +class ClusterType(Enum): + AWS = "aws" + FAIR = "fair" + RSC = "rsc" + + +def _guess_cluster_type() -> ClusterType: + uname = os.uname() + if uname.sysname == "Linux": + if uname.release.endswith("-aws"): + # Linux kernel versions on AWS instances are of the form "5.4.0-1051-aws" + return ClusterType.AWS + elif uname.nodename.startswith("rsc"): + # Linux kernel versions on RSC instances are standard ones but hostnames start with "rsc" + return ClusterType.RSC + + return ClusterType.FAIR + + +def get_cluster_type(cluster_type: Optional[ClusterType] = None) -> Optional[ClusterType]: + if cluster_type is None: + return _guess_cluster_type() + + return cluster_type + + +def get_checkpoint_path(cluster_type: Optional[ClusterType] = None) -> Optional[Path]: + cluster_type = get_cluster_type(cluster_type) + if cluster_type is None: + return None + + CHECKPOINT_DIRNAMES = { + ClusterType.AWS: "checkpoints", + ClusterType.FAIR: "checkpoint", + ClusterType.RSC: "checkpoint/dino", + } + return Path("/") / CHECKPOINT_DIRNAMES[cluster_type] + + +def get_user_checkpoint_path(cluster_type: Optional[ClusterType] = None) -> Optional[Path]: + checkpoint_path = get_checkpoint_path(cluster_type) + if checkpoint_path is None: + return None + + username = os.environ.get("USER") + assert username is not None + return checkpoint_path / username + + +def get_slurm_partition(cluster_type: Optional[ClusterType] = None) -> Optional[str]: + cluster_type = get_cluster_type(cluster_type) + if cluster_type is None: + return None + + SLURM_PARTITIONS = { + ClusterType.AWS: "learnlab", + ClusterType.FAIR: "learnlab", + ClusterType.RSC: "learn", + } + return SLURM_PARTITIONS[cluster_type] + + +def get_slurm_executor_parameters( + nodes: int, num_gpus_per_node: int, cluster_type: Optional[ClusterType] = None, **kwargs +) -> Dict[str, Any]: + # create default parameters + params = { + "mem_gb": 0, # Requests all memory on a node, see https://slurm.schedmd.com/sbatch.html + "gpus_per_node": num_gpus_per_node, + "tasks_per_node": num_gpus_per_node, # one task per GPU + "cpus_per_task": 10, + "nodes": nodes, + "slurm_partition": get_slurm_partition(cluster_type), + } + # apply cluster-specific adjustments + cluster_type = get_cluster_type(cluster_type) + if cluster_type == ClusterType.AWS: + params["cpus_per_task"] = 12 + del params["mem_gb"] + elif cluster_type == ClusterType.RSC: + params["cpus_per_task"] = 12 + # set additional parameters / apply overrides + params.update(kwargs) + return params diff --git a/moge/model/dinov2/utils/config.py b/moge/model/dinov2/utils/config.py new file mode 100644 index 0000000..c9de578 --- /dev/null +++ b/moge/model/dinov2/utils/config.py @@ -0,0 +1,72 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import math +import logging +import os + +from omegaconf import OmegaConf + +import dinov2.distributed as distributed +from dinov2.logging import setup_logging +from dinov2.utils import utils +from dinov2.configs import dinov2_default_config + + +logger = logging.getLogger("dinov2") + + +def apply_scaling_rules_to_cfg(cfg): # to fix + if cfg.optim.scaling_rule == "sqrt_wrt_1024": + base_lr = cfg.optim.base_lr + cfg.optim.lr = base_lr + cfg.optim.lr *= math.sqrt(cfg.train.batch_size_per_gpu * distributed.get_global_size() / 1024.0) + logger.info(f"sqrt scaling learning rate; base: {base_lr}, new: {cfg.optim.lr}") + else: + raise NotImplementedError + return cfg + + +def write_config(cfg, output_dir, name="config.yaml"): + logger.info(OmegaConf.to_yaml(cfg)) + saved_cfg_path = os.path.join(output_dir, name) + with open(saved_cfg_path, "w") as f: + OmegaConf.save(config=cfg, f=f) + return saved_cfg_path + + +def get_cfg_from_args(args): + args.output_dir = os.path.abspath(args.output_dir) + args.opts += [f"train.output_dir={args.output_dir}"] + default_cfg = OmegaConf.create(dinov2_default_config) + cfg = OmegaConf.load(args.config_file) + cfg = OmegaConf.merge(default_cfg, cfg, OmegaConf.from_cli(args.opts)) + return cfg + + +def default_setup(args): + distributed.enable(overwrite=True) + seed = getattr(args, "seed", 0) + rank = distributed.get_global_rank() + + global logger + setup_logging(output=args.output_dir, level=logging.INFO) + logger = logging.getLogger("dinov2") + + utils.fix_random_seeds(seed + rank) + logger.info("git:\n {}\n".format(utils.get_sha())) + logger.info("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items()))) + + +def setup(args): + """ + Create configs and perform basic setups. + """ + cfg = get_cfg_from_args(args) + os.makedirs(args.output_dir, exist_ok=True) + default_setup(args) + apply_scaling_rules_to_cfg(cfg) + write_config(cfg, args.output_dir) + return cfg diff --git a/moge/model/dinov2/utils/dtype.py b/moge/model/dinov2/utils/dtype.py new file mode 100644 index 0000000..80f4cd7 --- /dev/null +++ b/moge/model/dinov2/utils/dtype.py @@ -0,0 +1,37 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + + +from typing import Dict, Union + +import numpy as np +import torch + + +TypeSpec = Union[str, np.dtype, torch.dtype] + + +_NUMPY_TO_TORCH_DTYPE: Dict[np.dtype, torch.dtype] = { + np.dtype("bool"): torch.bool, + np.dtype("uint8"): torch.uint8, + np.dtype("int8"): torch.int8, + np.dtype("int16"): torch.int16, + np.dtype("int32"): torch.int32, + np.dtype("int64"): torch.int64, + np.dtype("float16"): torch.float16, + np.dtype("float32"): torch.float32, + np.dtype("float64"): torch.float64, + np.dtype("complex64"): torch.complex64, + np.dtype("complex128"): torch.complex128, +} + + +def as_torch_dtype(dtype: TypeSpec) -> torch.dtype: + if isinstance(dtype, torch.dtype): + return dtype + if isinstance(dtype, str): + dtype = np.dtype(dtype) + assert isinstance(dtype, np.dtype), f"Expected an instance of nunpy dtype, got {type(dtype)}" + return _NUMPY_TO_TORCH_DTYPE[dtype] diff --git a/moge/model/dinov2/utils/param_groups.py b/moge/model/dinov2/utils/param_groups.py new file mode 100644 index 0000000..9a5d2ff --- /dev/null +++ b/moge/model/dinov2/utils/param_groups.py @@ -0,0 +1,103 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +from collections import defaultdict +import logging + + +logger = logging.getLogger("dinov2") + + +def get_vit_lr_decay_rate(name, lr_decay_rate=1.0, num_layers=12, force_is_backbone=False, chunked_blocks=False): + """ + Calculate lr decay rate for different ViT blocks. + Args: + name (string): parameter name. + lr_decay_rate (float): base lr decay rate. + num_layers (int): number of ViT blocks. + Returns: + lr decay rate for the given parameter. + """ + layer_id = num_layers + 1 + if name.startswith("backbone") or force_is_backbone: + if ( + ".pos_embed" in name + or ".patch_embed" in name + or ".mask_token" in name + or ".cls_token" in name + or ".register_tokens" in name + ): + layer_id = 0 + elif force_is_backbone and ( + "pos_embed" in name + or "patch_embed" in name + or "mask_token" in name + or "cls_token" in name + or "register_tokens" in name + ): + layer_id = 0 + elif ".blocks." in name and ".residual." not in name: + layer_id = int(name[name.find(".blocks.") :].split(".")[2]) + 1 + elif chunked_blocks and "blocks." in name and "residual." not in name: + layer_id = int(name[name.find("blocks.") :].split(".")[2]) + 1 + elif "blocks." in name and "residual." not in name: + layer_id = int(name[name.find("blocks.") :].split(".")[1]) + 1 + + return lr_decay_rate ** (num_layers + 1 - layer_id) + + +def get_params_groups_with_decay(model, lr_decay_rate=1.0, patch_embed_lr_mult=1.0): + chunked_blocks = False + if hasattr(model, "n_blocks"): + logger.info("chunked fsdp") + n_blocks = model.n_blocks + chunked_blocks = model.chunked_blocks + elif hasattr(model, "blocks"): + logger.info("first code branch") + n_blocks = len(model.blocks) + elif hasattr(model, "backbone"): + logger.info("second code branch") + n_blocks = len(model.backbone.blocks) + else: + logger.info("else code branch") + n_blocks = 0 + all_param_groups = [] + + for name, param in model.named_parameters(): + name = name.replace("_fsdp_wrapped_module.", "") + if not param.requires_grad: + continue + decay_rate = get_vit_lr_decay_rate( + name, lr_decay_rate, num_layers=n_blocks, force_is_backbone=n_blocks > 0, chunked_blocks=chunked_blocks + ) + d = {"params": param, "is_last_layer": False, "lr_multiplier": decay_rate, "wd_multiplier": 1.0, "name": name} + + if "last_layer" in name: + d.update({"is_last_layer": True}) + + if name.endswith(".bias") or "norm" in name or "gamma" in name: + d.update({"wd_multiplier": 0.0}) + + if "patch_embed" in name: + d.update({"lr_multiplier": d["lr_multiplier"] * patch_embed_lr_mult}) + + all_param_groups.append(d) + logger.info(f"""{name}: lr_multiplier: {d["lr_multiplier"]}, wd_multiplier: {d["wd_multiplier"]}""") + + return all_param_groups + + +def fuse_params_groups(all_params_groups, keys=("lr_multiplier", "wd_multiplier", "is_last_layer")): + fused_params_groups = defaultdict(lambda: {"params": []}) + for d in all_params_groups: + identifier = "" + for k in keys: + identifier += k + str(d[k]) + "_" + + for k in keys: + fused_params_groups[identifier][k] = d[k] + fused_params_groups[identifier]["params"].append(d["params"]) + + return fused_params_groups.values() diff --git a/moge/model/dinov2/utils/utils.py b/moge/model/dinov2/utils/utils.py new file mode 100644 index 0000000..68f8e2c --- /dev/null +++ b/moge/model/dinov2/utils/utils.py @@ -0,0 +1,95 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import logging +import os +import random +import subprocess +from urllib.parse import urlparse + +import numpy as np +import torch +from torch import nn + + +logger = logging.getLogger("dinov2") + + +def load_pretrained_weights(model, pretrained_weights, checkpoint_key): + if urlparse(pretrained_weights).scheme: # If it looks like an URL + state_dict = torch.hub.load_state_dict_from_url(pretrained_weights, map_location="cpu") + else: + state_dict = torch.load(pretrained_weights, map_location="cpu") + if checkpoint_key is not None and checkpoint_key in state_dict: + logger.info(f"Take key {checkpoint_key} in provided checkpoint dict") + state_dict = state_dict[checkpoint_key] + # remove `module.` prefix + state_dict = {k.replace("module.", ""): v for k, v in state_dict.items()} + # remove `backbone.` prefix induced by multicrop wrapper + state_dict = {k.replace("backbone.", ""): v for k, v in state_dict.items()} + msg = model.load_state_dict(state_dict, strict=False) + logger.info("Pretrained weights found at {} and loaded with msg: {}".format(pretrained_weights, msg)) + + +def fix_random_seeds(seed=31): + """ + Fix random seeds. + """ + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + np.random.seed(seed) + random.seed(seed) + + +def get_sha(): + cwd = os.path.dirname(os.path.abspath(__file__)) + + def _run(command): + return subprocess.check_output(command, cwd=cwd).decode("ascii").strip() + + sha = "N/A" + diff = "clean" + branch = "N/A" + try: + sha = _run(["git", "rev-parse", "HEAD"]) + subprocess.check_output(["git", "diff"], cwd=cwd) + diff = _run(["git", "diff-index", "HEAD"]) + diff = "has uncommitted changes" if diff else "clean" + branch = _run(["git", "rev-parse", "--abbrev-ref", "HEAD"]) + except Exception: + pass + message = f"sha: {sha}, status: {diff}, branch: {branch}" + return message + + +class CosineScheduler(object): + def __init__(self, base_value, final_value, total_iters, warmup_iters=0, start_warmup_value=0, freeze_iters=0): + super().__init__() + self.final_value = final_value + self.total_iters = total_iters + + freeze_schedule = np.zeros((freeze_iters)) + + warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters) + + iters = np.arange(total_iters - warmup_iters - freeze_iters) + schedule = final_value + 0.5 * (base_value - final_value) * (1 + np.cos(np.pi * iters / len(iters))) + self.schedule = np.concatenate((freeze_schedule, warmup_schedule, schedule)) + + assert len(self.schedule) == self.total_iters + + def __getitem__(self, it): + if it >= self.total_iters: + return self.final_value + else: + return self.schedule[it] + + +def has_batchnorms(model): + bn_types = (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm) + for name, module in model.named_modules(): + if isinstance(module, bn_types): + return True + return False diff --git a/moge/model/modules.py b/moge/model/modules.py new file mode 100644 index 0000000..b36ad48 --- /dev/null +++ b/moge/model/modules.py @@ -0,0 +1,254 @@ +from typing import * +from numbers import Number +import importlib +import itertools +import functools +import sys + +import torch +from torch import Tensor +import torch.nn as nn +import torch.nn.functional as F + +from .dinov2.models.vision_transformer import DinoVisionTransformer +from .utils import wrap_dinov2_attention_with_sdpa, wrap_module_with_gradient_checkpointing, unwrap_module_with_gradient_checkpointing +from ..utils.geometry_torch import normalized_view_plane_uv + + +class ResidualConvBlock(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int = None, + hidden_channels: int = None, + kernel_size: int = 3, + padding_mode: str = 'replicate', + activation: Literal['relu', 'leaky_relu', 'silu', 'elu'] = 'relu', + in_norm: Literal['group_norm', 'layer_norm', 'instance_norm', 'none'] = 'layer_norm', + hidden_norm: Literal['group_norm', 'layer_norm', 'instance_norm'] = 'group_norm', + ): + super(ResidualConvBlock, self).__init__() + if out_channels is None: + out_channels = in_channels + if hidden_channels is None: + hidden_channels = in_channels + + if activation =='relu': + activation_cls = nn.ReLU + elif activation == 'leaky_relu': + activation_cls = functools.partial(nn.LeakyReLU, negative_slope=0.2) + elif activation =='silu': + activation_cls = nn.SiLU + elif activation == 'elu': + activation_cls = nn.ELU + else: + raise ValueError(f'Unsupported activation function: {activation}') + + self.layers = nn.Sequential( + nn.GroupNorm(in_channels // 32, in_channels) if in_norm == 'group_norm' else \ + nn.GroupNorm(1, in_channels) if in_norm == 'layer_norm' else \ + nn.InstanceNorm2d(in_channels) if in_norm == 'instance_norm' else \ + nn.Identity(), + activation_cls(), + nn.Conv2d(in_channels, hidden_channels, kernel_size=kernel_size, padding=kernel_size // 2, padding_mode=padding_mode), + nn.GroupNorm(hidden_channels // 32, hidden_channels) if hidden_norm == 'group_norm' else \ + nn.GroupNorm(1, hidden_channels) if hidden_norm == 'layer_norm' else \ + nn.InstanceNorm2d(hidden_channels) if hidden_norm == 'instance_norm' else\ + nn.Identity(), + activation_cls(), + nn.Conv2d(hidden_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2, padding_mode=padding_mode) + ) + + self.skip_connection = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0) if in_channels != out_channels else nn.Identity() + + def forward(self, x): + skip = self.skip_connection(x) + x = self.layers(x) + x = x + skip + return x + + +class DINOv2Encoder(nn.Module): + "Wrapped DINOv2 encoder supporting gradient checkpointing. Input is RGB image in range [0, 1]." + backbone: DinoVisionTransformer + image_mean: torch.Tensor + image_std: torch.Tensor + dim_features: int + + def __init__(self, backbone: str, intermediate_layers: Union[int, List[int]], dim_out: int, **deprecated_kwargs): + super(DINOv2Encoder, self).__init__() + + self.intermediate_layers = intermediate_layers + + # Load the backbone + self.hub_loader = getattr(importlib.import_module(".dinov2.hub.backbones", __package__), backbone) + self.backbone_name = backbone + self.backbone = self.hub_loader(pretrained=False) + + self.dim_features = self.backbone.blocks[0].attn.qkv.in_features + self.num_features = intermediate_layers if isinstance(intermediate_layers, int) else len(intermediate_layers) + + self.output_projections = nn.ModuleList([ + nn.Conv2d(in_channels=self.dim_features, out_channels=dim_out, kernel_size=1, stride=1, padding=0,) + for _ in range(self.num_features) + ]) + + self.register_buffer("image_mean", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)) + self.register_buffer("image_std", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)) + + @property + def onnx_compatible_mode(self): + return getattr(self, "_onnx_compatible_mode", False) + + @onnx_compatible_mode.setter + def onnx_compatible_mode(self, value: bool): + self._onnx_compatible_mode = value + self.backbone.onnx_compatible_mode = value + + def init_weights(self): + pretrained_backbone_state_dict = self.hub_loader(pretrained=True).state_dict() + self.backbone.load_state_dict(pretrained_backbone_state_dict) + + def enable_gradient_checkpointing(self): + for i in range(len(self.backbone.blocks)): + wrap_module_with_gradient_checkpointing(self.backbone.blocks[i]) + + def enable_pytorch_native_sdpa(self): + for i in range(len(self.backbone.blocks)): + wrap_dinov2_attention_with_sdpa(self.backbone.blocks[i].attn) + + def forward(self, image: torch.Tensor, token_rows: Union[int, torch.LongTensor], token_cols: Union[int, torch.LongTensor], return_class_token: bool = False) -> Tuple[torch.Tensor, torch.Tensor]: + image_14 = F.interpolate(image, (token_rows * 14, token_cols * 14), mode="bilinear", align_corners=False, antialias=not self.onnx_compatible_mode) + image_14 = (image_14 - self.image_mean) / self.image_std + + # Get intermediate layers from the backbone + features = self.backbone.get_intermediate_layers(image_14, n=self.intermediate_layers, return_class_token=True) + + # Project features to the desired dimensionality + x = torch.stack([ + proj(feat.permute(0, 2, 1).unflatten(2, (token_rows, token_cols)).contiguous()) + for proj, (feat, clstoken) in zip(self.output_projections, features) + ], dim=1).sum(dim=1) + + if return_class_token: + return x, features[-1][1] + else: + return x + + +class Resampler(nn.Sequential): + def __init__(self, + in_channels: int, + out_channels: int, + type_: Literal['pixel_shuffle', 'nearest', 'bilinear', 'conv_transpose', 'pixel_unshuffle', 'avg_pool', 'max_pool'], + scale_factor: int = 2, + ): + if type_ == 'pixel_shuffle': + nn.Sequential.__init__(self, + nn.Conv2d(in_channels, out_channels * (scale_factor ** 2), kernel_size=3, stride=1, padding=1, padding_mode='replicate'), + nn.PixelShuffle(scale_factor), + nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate') + ) + for i in range(1, scale_factor ** 2): + self[0].weight.data[i::scale_factor ** 2] = self[0].weight.data[0::scale_factor ** 2] + self[0].bias.data[i::scale_factor ** 2] = self[0].bias.data[0::scale_factor ** 2] + elif type_ in ['nearest', 'bilinear']: + nn.Sequential.__init__(self, + nn.Upsample(scale_factor=scale_factor, mode=type_, align_corners=False if type_ == 'bilinear' else None), + nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate') + ) + elif type_ == 'conv_transpose': + nn.Sequential.__init__(self, + nn.ConvTranspose2d(in_channels, out_channels, kernel_size=scale_factor, stride=scale_factor), + nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate') + ) + self[0].weight.data[:] = self[0].weight.data[:, :, :1, :1] + elif type_ == 'pixel_unshuffle': + nn.Sequential.__init__(self, + nn.PixelUnshuffle(scale_factor), + nn.Conv2d(in_channels * (scale_factor ** 2), out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate') + ) + elif type_ == 'avg_pool': + nn.Sequential.__init__(self, + nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate'), + nn.AvgPool2d(kernel_size=scale_factor, stride=scale_factor), + ) + elif type_ == 'max_pool': + nn.Sequential.__init__(self, + nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate'), + nn.MaxPool2d(kernel_size=scale_factor, stride=scale_factor), + ) + else: + raise ValueError(f'Unsupported resampler type: {type_}') + +class MLP(nn.Sequential): + def __init__(self, dims: Sequence[int]): + nn.Sequential.__init__(self, + *itertools.chain(*[ + (nn.Linear(dim_in, dim_out), nn.ReLU(inplace=True)) + for dim_in, dim_out in zip(dims[:-2], dims[1:-1]) + ]), + nn.Linear(dims[-2], dims[-1]), + ) + + +class ConvStack(nn.Module): + def __init__(self, + dim_in: List[Optional[int]], + dim_res_blocks: List[int], + dim_out: List[Optional[int]], + resamplers: Union[Literal['pixel_shuffle', 'nearest', 'bilinear', 'conv_transpose', 'pixel_unshuffle', 'avg_pool', 'max_pool'], List], + dim_times_res_block_hidden: int = 1, + num_res_blocks: int = 1, + res_block_in_norm: Literal['layer_norm', 'group_norm' , 'instance_norm', 'none'] = 'layer_norm', + res_block_hidden_norm: Literal['layer_norm', 'group_norm' , 'instance_norm', 'none'] = 'group_norm', + activation: Literal['relu', 'leaky_relu', 'silu', 'elu'] = 'relu', + ): + super().__init__() + self.input_blocks = nn.ModuleList([ + nn.Conv2d(dim_in_, dim_res_block_, kernel_size=1, stride=1, padding=0) if dim_in_ is not None else nn.Identity() + for dim_in_, dim_res_block_ in zip(dim_in if isinstance(dim_in, Sequence) else itertools.repeat(dim_in), dim_res_blocks) + ]) + self.resamplers = nn.ModuleList([ + Resampler(dim_prev, dim_succ, scale_factor=2, type_=resampler) + for i, (dim_prev, dim_succ, resampler) in enumerate(zip( + dim_res_blocks[:-1], + dim_res_blocks[1:], + resamplers if isinstance(resamplers, Sequence) else itertools.repeat(resamplers) + )) + ]) + self.res_blocks = nn.ModuleList([ + nn.Sequential( + *( + ResidualConvBlock( + dim_res_block_, dim_res_block_, dim_times_res_block_hidden * dim_res_block_, + activation=activation, in_norm=res_block_in_norm, hidden_norm=res_block_hidden_norm + ) for _ in range(num_res_blocks[i] if isinstance(num_res_blocks, list) else num_res_blocks) + ) + ) for i, dim_res_block_ in enumerate(dim_res_blocks) + ]) + self.output_blocks = nn.ModuleList([ + nn.Conv2d(dim_res_block_, dim_out_, kernel_size=1, stride=1, padding=0) if dim_out_ is not None else nn.Identity() + for dim_out_, dim_res_block_ in zip(dim_out if isinstance(dim_out, Sequence) else itertools.repeat(dim_out), dim_res_blocks) + ]) + + def enable_gradient_checkpointing(self): + for i in range(len(self.resamplers)): + self.resamplers[i] = wrap_module_with_gradient_checkpointing(self.resamplers[i]) + for i in range(len(self.res_blocks)): + for j in range(len(self.res_blocks[i])): + self.res_blocks[i][j] = wrap_module_with_gradient_checkpointing(self.res_blocks[i][j]) + + def forward(self, in_features: List[torch.Tensor]): + out_features = [] + for i in range(len(self.res_blocks)): + feature = self.input_blocks[i](in_features[i]) + if i == 0: + x = feature + elif feature is not None: + x = x + feature + x = self.res_blocks[i](x) + out_features.append(self.output_blocks[i](x)) + if i < len(self.res_blocks) - 1: + x = self.resamplers[i](x) + return out_features diff --git a/moge/model/utils.py b/moge/model/utils.py new file mode 100644 index 0000000..c50761d --- /dev/null +++ b/moge/model/utils.py @@ -0,0 +1,49 @@ +from typing import * + +import torch +import torch.nn as nn +import torch.nn.functional as F + +def wrap_module_with_gradient_checkpointing(module: nn.Module): + from torch.utils.checkpoint import checkpoint + class _CheckpointingWrapper(module.__class__): + _restore_cls = module.__class__ + def forward(self, *args, **kwargs): + return checkpoint(super().forward, *args, use_reentrant=False, **kwargs) + + module.__class__ = _CheckpointingWrapper + return module + + +def unwrap_module_with_gradient_checkpointing(module: nn.Module): + module.__class__ = module.__class__._restore_cls + + +def wrap_dinov2_attention_with_sdpa(module: nn.Module): + assert torch.__version__ >= '2.0', "SDPA requires PyTorch 2.0 or later" + class _AttentionWrapper(module.__class__): + def forward(self, x: torch.Tensor, attn_bias=None) -> torch.Tensor: + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) # (3, B, H, N, C // H) + + q, k, v = torch.unbind(qkv, 0) # (B, H, N, C // H) + + x = F.scaled_dot_product_attention(q, k, v, attn_bias) + x = x.permute(0, 2, 1, 3).reshape(B, N, C) + + x = self.proj(x) + x = self.proj_drop(x) + return x + module.__class__ = _AttentionWrapper + return module + + +def sync_ddp_hook(state, bucket: torch.distributed.GradBucket) -> torch.futures.Future[torch.Tensor]: + group_to_use = torch.distributed.group.WORLD + world_size = group_to_use.size() + grad = bucket.buffer() + grad.div_(world_size) + torch.distributed.all_reduce(grad, group=group_to_use) + fut = torch.futures.Future() + fut.set_result(grad) + return fut diff --git a/moge/model/v1.py b/moge/model/v1.py new file mode 100644 index 0000000..2513b86 --- /dev/null +++ b/moge/model/v1.py @@ -0,0 +1,392 @@ +from typing import * +from numbers import Number +from functools import partial +from pathlib import Path +import importlib +import warnings +import json + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.utils +import torch.utils.checkpoint +import torch.version +import utils3d +from huggingface_hub import hf_hub_download + + +from ..utils.geometry_torch import normalized_view_plane_uv, recover_focal_shift, gaussian_blur_2d, dilate_with_mask +from .utils import wrap_dinov2_attention_with_sdpa, wrap_module_with_gradient_checkpointing, unwrap_module_with_gradient_checkpointing +from ..utils.tools import timeit + + +class ResidualConvBlock(nn.Module): + def __init__(self, in_channels: int, out_channels: int = None, hidden_channels: int = None, padding_mode: str = 'replicate', activation: Literal['relu', 'leaky_relu', 'silu', 'elu'] = 'relu', norm: Literal['group_norm', 'layer_norm'] = 'group_norm'): + super(ResidualConvBlock, self).__init__() + if out_channels is None: + out_channels = in_channels + if hidden_channels is None: + hidden_channels = in_channels + + if activation =='relu': + activation_cls = lambda: nn.ReLU(inplace=True) + elif activation == 'leaky_relu': + activation_cls = lambda: nn.LeakyReLU(negative_slope=0.2, inplace=True) + elif activation =='silu': + activation_cls = lambda: nn.SiLU(inplace=True) + elif activation == 'elu': + activation_cls = lambda: nn.ELU(inplace=True) + else: + raise ValueError(f'Unsupported activation function: {activation}') + + self.layers = nn.Sequential( + nn.GroupNorm(1, in_channels), + activation_cls(), + nn.Conv2d(in_channels, hidden_channels, kernel_size=3, padding=1, padding_mode=padding_mode), + nn.GroupNorm(hidden_channels // 32 if norm == 'group_norm' else 1, hidden_channels), + activation_cls(), + nn.Conv2d(hidden_channels, out_channels, kernel_size=3, padding=1, padding_mode=padding_mode) + ) + + self.skip_connection = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0) if in_channels != out_channels else nn.Identity() + + def forward(self, x): + skip = self.skip_connection(x) + x = self.layers(x) + x = x + skip + return x + + +class Head(nn.Module): + def __init__( + self, + num_features: int, + dim_in: int, + dim_out: List[int], + dim_proj: int = 512, + dim_upsample: List[int] = [256, 128, 128], + dim_times_res_block_hidden: int = 1, + num_res_blocks: int = 1, + res_block_norm: Literal['group_norm', 'layer_norm'] = 'group_norm', + last_res_blocks: int = 0, + last_conv_channels: int = 32, + last_conv_size: int = 1 + ): + super().__init__() + + self.projects = nn.ModuleList([ + nn.Conv2d(in_channels=dim_in, out_channels=dim_proj, kernel_size=1, stride=1, padding=0,) for _ in range(num_features) + ]) + + self.upsample_blocks = nn.ModuleList([ + nn.Sequential( + self._make_upsampler(in_ch + 2, out_ch), + *(ResidualConvBlock(out_ch, out_ch, dim_times_res_block_hidden * out_ch, activation="relu", norm=res_block_norm) for _ in range(num_res_blocks)) + ) for in_ch, out_ch in zip([dim_proj] + dim_upsample[:-1], dim_upsample) + ]) + + self.output_block = nn.ModuleList([ + self._make_output_block( + dim_upsample[-1] + 2, dim_out_, dim_times_res_block_hidden, last_res_blocks, last_conv_channels, last_conv_size, res_block_norm, + ) for dim_out_ in dim_out + ]) + + def _make_upsampler(self, in_channels: int, out_channels: int): + upsampler = nn.Sequential( + nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2), + nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate') + ) + upsampler[0].weight.data[:] = upsampler[0].weight.data[:, :, :1, :1] + return upsampler + + def _make_output_block(self, dim_in: int, dim_out: int, dim_times_res_block_hidden: int, last_res_blocks: int, last_conv_channels: int, last_conv_size: int, res_block_norm: Literal['group_norm', 'layer_norm']): + return nn.Sequential( + nn.Conv2d(dim_in, last_conv_channels, kernel_size=3, stride=1, padding=1, padding_mode='replicate'), + *(ResidualConvBlock(last_conv_channels, last_conv_channels, dim_times_res_block_hidden * last_conv_channels, activation='relu', norm=res_block_norm) for _ in range(last_res_blocks)), + nn.ReLU(inplace=True), + nn.Conv2d(last_conv_channels, dim_out, kernel_size=last_conv_size, stride=1, padding=last_conv_size // 2, padding_mode='replicate'), + ) + + def forward(self, hidden_states: torch.Tensor, image: torch.Tensor): + img_h, img_w = image.shape[-2:] + patch_h, patch_w = img_h // 14, img_w // 14 + + # Process the hidden states + x = torch.stack([ + proj(feat.permute(0, 2, 1).unflatten(2, (patch_h, patch_w)).contiguous()) + for proj, (feat, clstoken) in zip(self.projects, hidden_states) + ], dim=1).sum(dim=1) + + # Upsample stage + # (patch_h, patch_w) -> (patch_h * 2, patch_w * 2) -> (patch_h * 4, patch_w * 4) -> (patch_h * 8, patch_w * 8) + for i, block in enumerate(self.upsample_blocks): + # UV coordinates is for awareness of image aspect ratio + uv = normalized_view_plane_uv(width=x.shape[-1], height=x.shape[-2], aspect_ratio=img_w / img_h, dtype=x.dtype, device=x.device) + uv = uv.permute(2, 0, 1).unsqueeze(0).expand(x.shape[0], -1, -1, -1) + x = torch.cat([x, uv], dim=1) + for layer in block: + x = torch.utils.checkpoint.checkpoint(layer, x, use_reentrant=False) + + # (patch_h * 8, patch_w * 8) -> (img_h, img_w) + x = F.interpolate(x, (img_h, img_w), mode="bilinear", align_corners=False) + uv = normalized_view_plane_uv(width=x.shape[-1], height=x.shape[-2], aspect_ratio=img_w / img_h, dtype=x.dtype, device=x.device) + uv = uv.permute(2, 0, 1).unsqueeze(0).expand(x.shape[0], -1, -1, -1) + x = torch.cat([x, uv], dim=1) + + if isinstance(self.output_block, nn.ModuleList): + output = [torch.utils.checkpoint.checkpoint(block, x, use_reentrant=False) for block in self.output_block] + else: + output = torch.utils.checkpoint.checkpoint(self.output_block, x, use_reentrant=False) + + return output + + +class MoGeModel(nn.Module): + image_mean: torch.Tensor + image_std: torch.Tensor + + def __init__(self, + encoder: str = 'dinov2_vitb14', + intermediate_layers: Union[int, List[int]] = 4, + dim_proj: int = 512, + dim_upsample: List[int] = [256, 128, 128], + dim_times_res_block_hidden: int = 1, + num_res_blocks: int = 1, + remap_output: Literal[False, True, 'linear', 'sinh', 'exp', 'sinh_exp'] = 'linear', + res_block_norm: Literal['group_norm', 'layer_norm'] = 'group_norm', + num_tokens_range: Tuple[Number, Number] = [1200, 2500], + last_res_blocks: int = 0, + last_conv_channels: int = 32, + last_conv_size: int = 1, + mask_threshold: float = 0.5, + **deprecated_kwargs + ): + super(MoGeModel, self).__init__() + + if deprecated_kwargs: + # Process legacy arguments + if 'trained_area_range' in deprecated_kwargs: + num_tokens_range = [deprecated_kwargs['trained_area_range'][0] // 14 ** 2, deprecated_kwargs['trained_area_range'][1] // 14 ** 2] + del deprecated_kwargs['trained_area_range'] + warnings.warn(f"The following deprecated/invalid arguments are ignored: {deprecated_kwargs}") + + self.encoder = encoder + self.remap_output = remap_output + self.intermediate_layers = intermediate_layers + self.num_tokens_range = num_tokens_range + self.mask_threshold = mask_threshold + + # NOTE: We have copied the DINOv2 code in torchhub to this repository. + # Minimal modifications have been made: removing irrelevant code, unnecessary warnings and fixing importing issues. + hub_loader = getattr(importlib.import_module(".dinov2.hub.backbones", __package__), encoder) + self.backbone = hub_loader(pretrained=False) + dim_feature = self.backbone.blocks[0].attn.qkv.in_features + + self.head = Head( + num_features=intermediate_layers if isinstance(intermediate_layers, int) else len(intermediate_layers), + dim_in=dim_feature, + dim_out=[3, 1], + dim_proj=dim_proj, + dim_upsample=dim_upsample, + dim_times_res_block_hidden=dim_times_res_block_hidden, + num_res_blocks=num_res_blocks, + res_block_norm=res_block_norm, + last_res_blocks=last_res_blocks, + last_conv_channels=last_conv_channels, + last_conv_size=last_conv_size + ) + + image_mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) + image_std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) + + self.register_buffer("image_mean", image_mean) + self.register_buffer("image_std", image_std) + + @property + def device(self) -> torch.device: + return next(self.parameters()).device + + @property + def dtype(self) -> torch.dtype: + return next(self.parameters()).dtype + + @classmethod + def from_pretrained(cls, pretrained_model_name_or_path: Union[str, Path, IO[bytes]], model_kwargs: Optional[Dict[str, Any]] = None, **hf_kwargs) -> 'MoGeModel': + """ + Load a model from a checkpoint file. + + ### Parameters: + - `pretrained_model_name_or_path`: path to the checkpoint file or repo id. + - `model_kwargs`: additional keyword arguments to override the parameters in the checkpoint. + - `hf_kwargs`: additional keyword arguments to pass to the `hf_hub_download` function. Ignored if `pretrained_model_name_or_path` is a local path. + + ### Returns: + - A new instance of `MoGe` with the parameters loaded from the checkpoint. + """ + if Path(pretrained_model_name_or_path).exists(): + checkpoint = torch.load(pretrained_model_name_or_path, map_location='cpu', weights_only=True) + else: + cached_checkpoint_path = hf_hub_download( + repo_id=pretrained_model_name_or_path, + repo_type="model", + filename="model.pt", + **hf_kwargs + ) + checkpoint = torch.load(cached_checkpoint_path, map_location='cpu', weights_only=True) + model_config = checkpoint['model_config'] + if model_kwargs is not None: + model_config.update(model_kwargs) + model = cls(**model_config) + model.load_state_dict(checkpoint['model']) + return model + + def init_weights(self): + "Load the backbone with pretrained dinov2 weights from torch hub" + state_dict = torch.hub.load('facebookresearch/dinov2', self.encoder, pretrained=True).state_dict() + self.backbone.load_state_dict(state_dict) + + def enable_gradient_checkpointing(self): + for i in range(len(self.backbone.blocks)): + self.backbone.blocks[i] = wrap_module_with_gradient_checkpointing(self.backbone.blocks[i]) + + def _remap_points(self, points: torch.Tensor) -> torch.Tensor: + if self.remap_output == 'linear': + pass + elif self.remap_output =='sinh': + points = torch.sinh(points) + elif self.remap_output == 'exp': + xy, z = points.split([2, 1], dim=-1) + z = torch.exp(z) + points = torch.cat([xy * z, z], dim=-1) + elif self.remap_output =='sinh_exp': + xy, z = points.split([2, 1], dim=-1) + points = torch.cat([torch.sinh(xy), torch.exp(z)], dim=-1) + else: + raise ValueError(f"Invalid remap output type: {self.remap_output}") + return points + + def forward(self, image: torch.Tensor, num_tokens: int) -> Dict[str, torch.Tensor]: + original_height, original_width = image.shape[-2:] + + # Resize to expected resolution defined by num_tokens + resize_factor = ((num_tokens * 14 ** 2) / (original_height * original_width)) ** 0.5 + resized_width, resized_height = int(original_width * resize_factor), int(original_height * resize_factor) + image = F.interpolate(image, (resized_height, resized_width), mode="bicubic", align_corners=False, antialias=True) + + # Apply image transformation for DINOv2 + image = (image - self.image_mean) / self.image_std + image_14 = F.interpolate(image, (resized_height // 14 * 14, resized_width // 14 * 14), mode="bilinear", align_corners=False, antialias=True) + + # Get intermediate layers from the backbone + features = self.backbone.get_intermediate_layers(image_14, self.intermediate_layers, return_class_token=True) + + # Predict points (and mask) + output = self.head(features, image) + points, mask = output + + # Make sure fp32 precision for output + with torch.autocast(device_type=image.device.type, dtype=torch.float32): + # Resize to original resolution + points = F.interpolate(points, (original_height, original_width), mode='bilinear', align_corners=False, antialias=False) + mask = F.interpolate(mask, (original_height, original_width), mode='bilinear', align_corners=False, antialias=False) + + # Post-process points and mask + points, mask = points.permute(0, 2, 3, 1), mask.squeeze(1) + points = self._remap_points(points) # slightly improves the performance in case of very large output values + + return_dict = {'points': points, 'mask': mask} + return return_dict + + @torch.inference_mode() + def infer( + self, + image: torch.Tensor, + fov_x: Union[Number, torch.Tensor] = None, + resolution_level: int = 9, + num_tokens: int = None, + apply_mask: bool = True, + force_projection: bool = True, + use_fp16: bool = True, + ) -> Dict[str, torch.Tensor]: + """ + User-friendly inference function + + ### Parameters + - `image`: input image tensor of shape (B, 3, H, W) or (3, H, W)\ + - `fov_x`: the horizontal camera FoV in degrees. If None, it will be inferred from the predicted point map. Default: None + - `resolution_level`: An integer [0-9] for the resolution level for inference. + The higher, the finer details will be captured, but slower. Defaults to 9. Note that it is irrelevant to the output size, which is always the same as the input size. + `resolution_level` actually controls `num_tokens`. See `num_tokens` for more details. + - `num_tokens`: number of tokens used for inference. A integer in the (suggested) range of `[1200, 2500]`. + `resolution_level` will be ignored if `num_tokens` is provided. Default: None + - `apply_mask`: if True, the output point map will be masked using the predicted mask. Default: True + - `force_projection`: if True, the output point map will be recomputed to match the projection constraint. Default: True + - `use_fp16`: if True, use mixed precision to speed up inference. Default: True + + ### Returns + + A dictionary containing the following keys: + - `points`: output tensor of shape (B, H, W, 3) or (H, W, 3). + - `depth`: tensor of shape (B, H, W) or (H, W) containing the depth map. + - `intrinsics`: tensor of shape (B, 3, 3) or (3, 3) containing the camera intrinsics. + """ + if image.dim() == 3: + omit_batch_dim = True + image = image.unsqueeze(0) + else: + omit_batch_dim = False + image = image.to(dtype=self.dtype, device=self.device) + + original_height, original_width = image.shape[-2:] + aspect_ratio = original_width / original_height + + if num_tokens is None: + min_tokens, max_tokens = self.num_tokens_range + num_tokens = int(min_tokens + (resolution_level / 9) * (max_tokens - min_tokens)) + + with torch.autocast(device_type=self.device.type, dtype=torch.float16, enabled=use_fp16 and self.dtype != torch.float16): + output = self.forward(image, num_tokens) + points, mask = output['points'], output['mask'] + + # Always process the output in fp32 precision + with torch.autocast(device_type=self.device.type, dtype=torch.float32): + points, mask, fov_x = map(lambda x: x.float() if isinstance(x, torch.Tensor) else x, [points, mask, fov_x]) + + mask_binary = mask > self.mask_threshold + + # Get camera-space point map. (Focal here is the focal length relative to half the image diagonal) + if fov_x is None: + focal, shift = recover_focal_shift(points, mask_binary) + else: + focal = aspect_ratio / (1 + aspect_ratio ** 2) ** 0.5 / torch.tan(torch.deg2rad(torch.as_tensor(fov_x, device=points.device, dtype=points.dtype) / 2)) + if focal.ndim == 0: + focal = focal[None].expand(points.shape[0]) + _, shift = recover_focal_shift(points, mask_binary, focal=focal) + fx = focal / 2 * (1 + aspect_ratio ** 2) ** 0.5 / aspect_ratio + fy = focal / 2 * (1 + aspect_ratio ** 2) ** 0.5 + intrinsics = utils3d.pt.intrinsics_from_focal_center(fx, fy, torch.tensor(0.5, device=points.device, dtype=points.dtype), torch.tensor(0.5, device=points.device, dtype=points.dtype)) + depth = points[..., 2] + shift[..., None, None] + + # If projection constraint is forced, recompute the point map using the actual depth map + if force_projection: + points = utils3d.pt.depth_map_to_point_map(depth, intrinsics=intrinsics) + else: + points = points + torch.stack([torch.zeros_like(shift), torch.zeros_like(shift), shift], dim=-1)[..., None, None, :] + + # Apply mask if needed + if apply_mask: + points = torch.where(mask_binary[..., None], points, torch.inf) + depth = torch.where(mask_binary, depth, torch.inf) + + return_dict = { + 'points': points, + 'intrinsics': intrinsics, + 'depth': depth, + 'mask': mask_binary, + } + + if omit_batch_dim: + return_dict = {k: v.squeeze(0) for k, v in return_dict.items()} + + return return_dict \ No newline at end of file diff --git a/moge/model/v2.py b/moge/model/v2.py new file mode 100644 index 0000000..5cf8028 --- /dev/null +++ b/moge/model/v2.py @@ -0,0 +1,303 @@ +from typing import * +from numbers import Number +from functools import partial +from pathlib import Path +import warnings + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.utils +import torch.utils.checkpoint +import torch.amp +import torch.version +import utils3d +from huggingface_hub import hf_hub_download + +from ..utils.geometry_torch import normalized_view_plane_uv, recover_focal_shift, angle_diff_vec3 +from .utils import wrap_dinov2_attention_with_sdpa, wrap_module_with_gradient_checkpointing, unwrap_module_with_gradient_checkpointing +from .modules import DINOv2Encoder, MLP, ConvStack + + +class MoGeModel(nn.Module): + encoder: DINOv2Encoder + neck: ConvStack + points_head: ConvStack + mask_head: ConvStack + scale_head: MLP + onnx_compatible_mode: bool + + def __init__(self, + encoder: Dict[str, Any], + neck: Dict[str, Any], + points_head: Dict[str, Any] = None, + mask_head: Dict[str, Any] = None, + normal_head: Dict[str, Any] = None, + scale_head: Dict[str, Any] = None, + remap_output: Literal['linear', 'sinh', 'exp', 'sinh_exp'] = 'linear', + num_tokens_range: List[int] = [1200, 3600], + **deprecated_kwargs + ): + super(MoGeModel, self).__init__() + if deprecated_kwargs: + warnings.warn(f"The following deprecated/invalid arguments are ignored: {deprecated_kwargs}") + + self.remap_output = remap_output + self.num_tokens_range = num_tokens_range + + self.encoder = DINOv2Encoder(**encoder) + self.neck = ConvStack(**neck) + if points_head is not None: + self.points_head = ConvStack(**points_head) + if mask_head is not None: + self.mask_head = ConvStack(**mask_head) + if normal_head is not None: + self.normal_head = ConvStack(**normal_head) + if scale_head is not None: + self.scale_head = MLP(**scale_head) + + @property + def device(self) -> torch.device: + return next(self.parameters()).device + + @property + def dtype(self) -> torch.dtype: + return next(self.parameters()).dtype + + @property + def onnx_compatible_mode(self) -> bool: + return getattr(self, "_onnx_compatible_mode", False) + + @onnx_compatible_mode.setter + def onnx_compatible_mode(self, value: bool): + self._onnx_compatible_mode = value + self.encoder.onnx_compatible_mode = value + + @classmethod + def from_pretrained(cls, pretrained_model_name_or_path: Union[str, Path, IO[bytes]], model_kwargs: Optional[Dict[str, Any]] = None, **hf_kwargs) -> 'MoGeModel': + """ + Load a model from a checkpoint file. + + ### Parameters: + - `pretrained_model_name_or_path`: path to the checkpoint file or repo id. + - `compiled` + - `model_kwargs`: additional keyword arguments to override the parameters in the checkpoint. + - `hf_kwargs`: additional keyword arguments to pass to the `hf_hub_download` function. Ignored if `pretrained_model_name_or_path` is a local path. + + ### Returns: + - A new instance of `MoGe` with the parameters loaded from the checkpoint. + """ + if Path(pretrained_model_name_or_path).exists(): + checkpoint_path = pretrained_model_name_or_path + else: + checkpoint_path = hf_hub_download( + repo_id=pretrained_model_name_or_path, + repo_type="model", + filename="model.pt", + **hf_kwargs + ) + checkpoint = torch.load(checkpoint_path, map_location='cpu', weights_only=True) + + model_config = checkpoint['model_config'] + if model_kwargs is not None: + model_config.update(model_kwargs) + model = cls(**model_config) + model.load_state_dict(checkpoint['model'], strict=False) + + return model + + def init_weights(self): + self.encoder.init_weights() + + def enable_gradient_checkpointing(self): + self.encoder.enable_gradient_checkpointing() + self.neck.enable_gradient_checkpointing() + for head in ['points_head', 'normal_head', 'mask_head']: + if hasattr(self, head): + getattr(self, head).enable_gradient_checkpointing() + + def enable_pytorch_native_sdpa(self): + self.encoder.enable_pytorch_native_sdpa() + + def _remap_points(self, points: torch.Tensor) -> torch.Tensor: + if self.remap_output == 'linear': + pass + elif self.remap_output =='sinh': + points = torch.sinh(points) + elif self.remap_output == 'exp': + xy, z = points.split([2, 1], dim=-1) + z = torch.exp(z) + points = torch.cat([xy * z, z], dim=-1) + elif self.remap_output =='sinh_exp': + xy, z = points.split([2, 1], dim=-1) + points = torch.cat([torch.sinh(xy), torch.exp(z)], dim=-1) + else: + raise ValueError(f"Invalid remap output type: {self.remap_output}") + return points + + def forward(self, image: torch.Tensor, num_tokens: Union[int, torch.LongTensor]) -> Dict[str, torch.Tensor]: + batch_size, _, img_h, img_w = image.shape + device, dtype = image.device, image.dtype + + aspect_ratio = img_w / img_h + base_h, base_w = (num_tokens / aspect_ratio) ** 0.5, (num_tokens * aspect_ratio) ** 0.5 + if isinstance(base_h, torch.Tensor): + base_h, base_w = base_h.round().long(), base_w.round().long() + else: + base_h, base_w = round(base_h), round(base_w) + + # Backbones encoding + features, cls_token = self.encoder(image, base_h, base_w, return_class_token=True) + features = [features, None, None, None, None] + + # Concat UVs for aspect ratio input + for level in range(5): + uv = normalized_view_plane_uv(width=base_w * 2 ** level, height=base_h * 2 ** level, aspect_ratio=aspect_ratio, dtype=dtype, device=device) + uv = uv.permute(2, 0, 1).unsqueeze(0).expand(batch_size, -1, -1, -1) + if features[level] is None: + features[level] = uv + else: + features[level] = torch.concat([features[level], uv], dim=1) + + # Shared neck + features = self.neck(features) + + # Heads decoding + points, normal, mask = (getattr(self, head)(features)[-1] if hasattr(self, head) else None for head in ['points_head', 'normal_head', 'mask_head']) + metric_scale = self.scale_head(cls_token) if hasattr(self, 'scale_head') else None + + # Resize + points, normal, mask = (F.interpolate(v, (img_h, img_w), mode='bilinear', align_corners=False, antialias=False) if v is not None else None for v in [points, normal, mask]) + + # Remap output + if points is not None: + points = points.permute(0, 2, 3, 1) + points = self._remap_points(points) # slightly improves the performance in case of very large output values + if normal is not None: + normal = normal.permute(0, 2, 3, 1) + normal = F.normalize(normal, dim=-1) + if mask is not None: + mask = mask.squeeze(1).sigmoid() + if metric_scale is not None: + metric_scale = metric_scale.squeeze(1).exp() + + return_dict = { + 'points': points, + 'normal': normal, + 'mask': mask, + 'metric_scale': metric_scale + } + return_dict = {k: v for k, v in return_dict.items() if v is not None} + + return return_dict + + @torch.inference_mode() + def infer( + self, + image: torch.Tensor, + num_tokens: int = None, + resolution_level: int = 9, + force_projection: bool = True, + apply_mask: bool = True, + fov_x: Optional[Union[Number, torch.Tensor]] = None, + use_fp16: bool = True, + ) -> Dict[str, torch.Tensor]: + """ + User-friendly inference function + + ### Parameters + - `image`: input image tensor of shape (B, 3, H, W) or (3, H, W) + - `num_tokens`: the number of base ViT tokens to use for inference, `'least'` or `'most'` or an integer. Suggested range: 1200 ~ 2500. + More tokens will result in significantly higher accuracy and finer details, but slower inference time. Default: `'most'`. + - `force_projection`: if True, the output point map will be computed using the actual depth map. Default: True + - `apply_mask`: if True, the output point map will be masked using the predicted mask. Default: True + - `fov_x`: the horizontal camera FoV in degrees. If None, it will be inferred from the predicted point map. Default: None + - `use_fp16`: if True, use mixed precision to speed up inference. Default: True + + ### Returns + + A dictionary containing the following keys: + - `points`: output tensor of shape (B, H, W, 3) or (H, W, 3). + - `depth`: tensor of shape (B, H, W) or (H, W) containing the depth map. + - `intrinsics`: tensor of shape (B, 3, 3) or (3, 3) containing the camera intrinsics. + """ + if image.dim() == 3: + omit_batch_dim = True + image = image.unsqueeze(0) + else: + omit_batch_dim = False + image = image.to(dtype=self.dtype, device=self.device) + + original_height, original_width = image.shape[-2:] + area = original_height * original_width + aspect_ratio = original_width / original_height + + # Determine the number of base tokens to use + if num_tokens is None: + min_tokens, max_tokens = self.num_tokens_range + num_tokens = int(min_tokens + (resolution_level / 9) * (max_tokens - min_tokens)) + + # Forward pass + with torch.autocast(device_type=self.device.type, dtype=torch.float16, enabled=use_fp16 and self.dtype != torch.float16): + output = self.forward(image, num_tokens=num_tokens) + points, normal, mask, metric_scale = (output.get(k, None) for k in ['points', 'normal', 'mask', 'metric_scale']) + + # Always process the output in fp32 precision + points, normal, mask, metric_scale, fov_x = map(lambda x: x.float() if isinstance(x, torch.Tensor) else x, [points, normal, mask, metric_scale, fov_x]) + with torch.autocast(device_type=self.device.type, dtype=torch.float32): + if mask is not None: + mask_binary = mask > 0.5 + else: + mask_binary = None + + if points is not None: + # Convert affine point map to camera-space. Recover depth and intrinsics from point map. + # NOTE: Focal here is the focal length relative to half the image diagonal + if fov_x is None: + # Recover focal and shift from predicted point map + focal, shift = recover_focal_shift(points, mask_binary) + else: + # Focal is known, recover shift only + focal = aspect_ratio / (1 + aspect_ratio ** 2) ** 0.5 / torch.tan(torch.deg2rad(torch.as_tensor(fov_x, device=points.device, dtype=points.dtype) / 2)) + if focal.ndim == 0: + focal = focal[None].expand(points.shape[0]) + _, shift = recover_focal_shift(points, mask_binary, focal=focal) + fx, fy = focal / 2 * (1 + aspect_ratio ** 2) ** 0.5 / aspect_ratio, focal / 2 * (1 + aspect_ratio ** 2) ** 0.5 + intrinsics = utils3d.pt.intrinsics_from_focal_center(fx, fy, torch.tensor(0.5, device=points.device, dtype=points.dtype), torch.tensor(0.5, device=points.device, dtype=points.dtype)) + points[..., 2] += shift[..., None, None] + if mask_binary is not None: + mask_binary &= points[..., 2] > 0 # in case depth is contains negative values (which should never happen in practice) + depth = points[..., 2].clone() + else: + depth, intrinsics = None, None + + # If projection constraint is forced, recompute the point map using the actual depth map & intrinsics + if force_projection and depth is not None: + points = utils3d.pt.depth_map_to_point_map(depth, intrinsics=intrinsics) + + # Apply metric scale + if metric_scale is not None: + if points is not None: + points *= metric_scale[:, None, None, None] + if depth is not None: + depth *= metric_scale[:, None, None] + + # Apply mask + if apply_mask and mask_binary is not None: + points = torch.where(mask_binary[..., None], points, torch.inf) if points is not None else None + depth = torch.where(mask_binary, depth, torch.inf) if depth is not None else None + normal = torch.where(mask_binary[..., None], normal, torch.zeros_like(normal)) if normal is not None else None + + return_dict = { + 'points': points, + 'intrinsics': intrinsics, + 'depth': depth, + 'mask': mask_binary, + 'normal': normal + } + return_dict = {k: v for k, v in return_dict.items() if v is not None} + + if omit_batch_dim: + return_dict = {k: v.squeeze(0) for k, v in return_dict.items()} + + return return_dict diff --git a/moge/scripts/__init__.py b/moge/scripts/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/moge/scripts/app.py b/moge/scripts/app.py new file mode 100644 index 0000000..9a63e62 --- /dev/null +++ b/moge/scripts/app.py @@ -0,0 +1,301 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +import sys +from pathlib import Path +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) +import time +import uuid +import tempfile +import itertools +from typing import * +import atexit +from concurrent.futures import ThreadPoolExecutor +import shutil + +import click + + +@click.command(help='Web demo') +@click.option('--share', is_flag=True, help='Whether to run the app in shared mode.') +@click.option('--pretrained', 'pretrained_model_name_or_path', default=None, help='The name or path of the pre-trained model.') +@click.option('--version', 'model_version', default='v2', help='The version of the model.') +@click.option('--fp16', 'use_fp16', is_flag=True, help='Whether to use fp16 inference.') +def main(share: bool, pretrained_model_name_or_path: str, model_version: str, use_fp16: bool): + print("Import modules...") + # Lazy import + import cv2 + import torch + import numpy as np + import trimesh + import trimesh.visual + from PIL import Image + import gradio as gr + try: + import spaces # This is for deployment at huggingface.co/spaces + HUGGINFACE_SPACES_INSTALLED = True + except ImportError: + HUGGINFACE_SPACES_INSTALLED = False + + import utils3d + from moge.utils.io import write_normal + from moge.utils.vis import colorize_depth, colorize_normal + from moge.model import import_model_class_by_version + from moge.utils.geometry_numpy import depth_occlusion_edge_numpy + from moge.utils.tools import timeit + + print("Load model...") + if pretrained_model_name_or_path is None: + DEFAULT_PRETRAINED_MODEL_FOR_EACH_VERSION = { + "v1": "Ruicheng/moge-vitl", + "v2": "Ruicheng/moge-2-vitl-normal", + } + pretrained_model_name_or_path = DEFAULT_PRETRAINED_MODEL_FOR_EACH_VERSION[model_version] + model = import_model_class_by_version(model_version).from_pretrained(pretrained_model_name_or_path).cuda().eval() + if use_fp16: + model.half() + thread_pool_executor = ThreadPoolExecutor(max_workers=1) + + def delete_later(path: Union[str, os.PathLike], delay: int = 300): + def _delete(): + try: + os.remove(path) + except FileNotFoundError: + pass + def _wait_and_delete(): + time.sleep(delay) + _delete(path) + thread_pool_executor.submit(_wait_and_delete) + atexit.register(_delete) + + # Inference on GPU. + @(spaces.GPU if HUGGINFACE_SPACES_INSTALLED else lambda x: x) + def run_with_gpu(image: np.ndarray, resolution_level: int, apply_mask: bool) -> Dict[str, np.ndarray]: + image_tensor = torch.tensor(image, dtype=torch.float32 if not use_fp16 else torch.float16, device=torch.device('cuda')).permute(2, 0, 1) / 255 + output = model.infer(image_tensor, apply_mask=apply_mask, resolution_level=resolution_level, use_fp16=use_fp16) + output = {k: v.cpu().numpy() for k, v in output.items()} + return output + + # Full inference pipeline + def run(image: np.ndarray, max_size: int = 800, resolution_level: str = 'High', apply_mask: bool = True, remove_edge: bool = True, request: gr.Request = None): + larger_size = max(image.shape[:2]) + if larger_size > max_size: + scale = max_size / larger_size + image = cv2.resize(image, (0, 0), fx=scale, fy=scale, interpolation=cv2.INTER_AREA) + + height, width = image.shape[:2] + + resolution_level_int = {'Low': 0, 'Medium': 5, 'High': 9, 'Ultra': 30}.get(resolution_level, 9) + output = run_with_gpu(image, resolution_level_int, apply_mask) + + points, depth, mask, normal = output['points'], output['depth'], output['mask'], output.get('normal', None) + + if remove_edge: + mask_cleaned = mask & ~utils3d.np.depth_map_edge(depth, rtol=0.04) + else: + mask_cleaned = mask + + results = { + **output, + 'mask_cleaned': mask_cleaned, + 'image': image + } + + # depth & normal visualization + depth_vis = colorize_depth(depth) + if normal is not None: + normal_vis = colorize_normal(normal) + else: + normal_vis = gr.update(label="Normal map (not avalable for this model)") + + # mesh & pointcloud + if normal is None: + faces, vertices, vertex_colors, vertex_uvs = utils3d.np.build_mesh_from_map( + points, + image.astype(np.float32) / 255, + utils3d.np.uv_map(height, width), + mask=mask_cleaned, + tri=True + ) + vertex_normals = None + else: + faces, vertices, vertex_colors, vertex_uvs, vertex_normals = utils3d.np.build_mesh_from_map( + points, + image.astype(np.float32) / 255, + utils3d.np.uv_map(height, width), + normal, + mask=mask_cleaned, + tri=True + ) + vertices = vertices * np.array([1, -1, -1], dtype=np.float32) + vertex_uvs = vertex_uvs * np.array([1, -1], dtype=np.float32) + np.array([0, 1], dtype=np.float32) + if vertex_normals is not None: + vertex_normals = vertex_normals * np.array([1, -1, -1], dtype=np.float32) + + tempdir = Path(tempfile.gettempdir(), 'moge') + tempdir.mkdir(exist_ok=True) + output_path = Path(tempdir, request.session_hash) + shutil.rmtree(output_path, ignore_errors=True) + output_path.mkdir(exist_ok=True, parents=True) + trimesh.Trimesh( + vertices=vertices, + faces=faces, + visual = trimesh.visual.texture.TextureVisuals( + uv=vertex_uvs, + material=trimesh.visual.material.PBRMaterial( + baseColorTexture=Image.fromarray(image), + metallicFactor=0.5, + roughnessFactor=1.0 + ) + ), + vertex_normals=vertex_normals, + process=False + ).export(output_path / 'mesh.glb') + pointcloud = trimesh.PointCloud( + vertices=vertices, + colors=vertex_colors, + ) + pointcloud.vertex_normals = vertex_normals + pointcloud.export(output_path / 'pointcloud.ply', vertex_normal=True) + trimesh.PointCloud( + vertices=vertices, + colors=vertex_colors, + ).export(output_path / 'pointcloud.glb', include_normals=True) + cv2.imwrite(str(output_path /'mask.png'), mask.astype(np.uint8) * 255) + cv2.imwrite(str(output_path / 'depth.exr'), depth.astype(np.float32), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + cv2.imwrite(str(output_path / 'points.exr'), cv2.cvtColor(points.astype(np.float32), cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + if normal is not None: + cv2.imwrite(str(output_path / 'normal.exr'), cv2.cvtColor(normal.astype(np.float32) * np.array([1, -1, -1], dtype=np.float32), cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_HALF]) + + files = ['mesh.glb', 'pointcloud.ply', 'depth.exr', 'points.exr', 'mask.png'] + if normal is not None: + files.append('normal.exr') + + for f in files: + delete_later(output_path / f) + + # FOV + intrinsics = results['intrinsics'] + fov_x, fov_y = utils3d.np.intrinsics_to_fov(intrinsics) + fov_x, fov_y = np.rad2deg([fov_x, fov_y]) + + # messages + viewer_message = f'**Note:** Inference has been completed. It may take a few seconds to download the 3D model.' + if resolution_level != 'Ultra': + depth_message = f'**Note:** Want sharper depth map? Try increasing the `maximum image size` and setting the `inference resolution level` to `Ultra` in the settings.' + else: + depth_message = "" + + return ( + results, + depth_vis, + normal_vis, + output_path / 'pointcloud.glb', + [(output_path / f).as_posix() for f in files if (output_path / f).exists()], + f'- **Horizontal FOV: {fov_x:.1f}°**. \n - **Vertical FOV: {fov_y:.1f}°**', + viewer_message, + depth_message + ) + + def reset_measure(results: Dict[str, np.ndarray]): + return [results['image'], [], ""] + + + def measure(results: Dict[str, np.ndarray], measure_points: List[Tuple[int, int]], event: gr.SelectData): + point2d = event.index[0], event.index[1] + measure_points.append(point2d) + + image = results['image'].copy() + for p in measure_points: + image = cv2.circle(image, p, radius=5, color=(255, 0, 0), thickness=2) + + depth_text = "" + for i, p in enumerate(measure_points): + d = results['depth'][p[1], p[0]] + depth_text += f"- **P{i + 1} depth: {d:.2f}m.**\n" + + if len(measure_points) == 2: + point1, point2 = measure_points + image = cv2.line(image, point1, point2, color=(255, 0, 0), thickness=2) + distance = np.linalg.norm(results['points'][point1[1], point1[0]] - results['points'][point2[1], point2[0]]) + measure_points = [] + + distance_text = f"- **Distance: {distance:.2f}m**" + + text = depth_text + distance_text + return [image, measure_points, text] + else: + return [image, measure_points, depth_text] + + print("Create Gradio app...") + with gr.Blocks(theme=gr.themes.Soft()) as demo: + gr.Markdown( +f''' +
+

Turn a 2D image into 3D with MoGe badge-github-stars

+
+''') + results = gr.State(value=None) + measure_points = gr.State(value=[]) + + with gr.Row(): + with gr.Column(): + input_image = gr.Image(type="numpy", image_mode="RGB", label="Input Image") + with gr.Accordion(label="Settings", open=False): + max_size_input = gr.Number(value=800, label="Maximum Image Size", precision=0, minimum=256, maximum=2048) + resolution_level = gr.Dropdown(['Low', 'Medium', 'High', 'Ultra'], label="Inference Resolution Level", value='High') + apply_mask = gr.Checkbox(value=True, label="Apply mask") + remove_edges = gr.Checkbox(value=True, label="Remove edges") + submit_btn = gr.Button("Submit", variant='primary') + + with gr.Column(): + with gr.Tabs(): + with gr.Tab("3D View"): + viewer_message = gr.Markdown("") + model_3d = gr.Model3D(display_mode="solid", label="3D Point Map", clear_color=[1.0, 1.0, 1.0, 1.0], height="60vh") + fov = gr.Markdown() + with gr.Tab("Depth"): + depth_message = gr.Markdown("") + depth_map = gr.Image(type="numpy", label="Colorized Depth Map", format='png', interactive=False) + with gr.Tab("Normal", interactive=hasattr(model, 'normal_head')): + normal_map = gr.Image(type="numpy", label="Normal Map", format='png', interactive=False) + with gr.Tab("Measure", interactive=hasattr(model, 'scale_head')): + gr.Markdown("### Click on the image to measure the distance between two points. \n" + "**Note:** Metric scale is most reliable for typical indoor or street scenes, and may degrade for contents unfamiliar to the model (e.g., stylized or close-up images).") + measure_image = gr.Image(type="numpy", show_label=False, format='webp', interactive=False, sources=[]) + measure_text = gr.Markdown("") + with gr.Tab("Download"): + files = gr.File(type='filepath', label="Output Files") + + if Path('example_images').exists(): + example_image_paths = sorted(list(itertools.chain(*[Path('example_images').glob(f'*.{ext}') for ext in ['jpg', 'png', 'jpeg', 'JPG', 'PNG', 'JPEG']]))) + examples = gr.Examples( + examples = example_image_paths, + inputs=input_image, + label="Examples" + ) + + submit_btn.click( + fn=lambda: [None, None, None, None, None, "", "", ""], + outputs=[results, depth_map, normal_map, model_3d, files, fov, viewer_message, depth_message] + ).then( + fn=run, + inputs=[input_image, max_size_input, resolution_level, apply_mask, remove_edges], + outputs=[results, depth_map, normal_map, model_3d, files, fov, viewer_message, depth_message] + ).then( + fn=reset_measure, + inputs=[results], + outputs=[measure_image, measure_points, measure_text] + ) + + measure_image.select( + fn=measure, + inputs=[results, measure_points], + outputs=[measure_image, measure_points, measure_text] + ) + + demo.launch(share=share) + + +if __name__ == '__main__': + main() diff --git a/moge/scripts/cli.py b/moge/scripts/cli.py new file mode 100644 index 0000000..45c3b90 --- /dev/null +++ b/moge/scripts/cli.py @@ -0,0 +1,27 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +from pathlib import Path +import sys +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) + +import click + + +@click.group(help='MoGe command line interface.') +def cli(): + pass + +def main(): + from moge.scripts import app, infer, infer_baseline, infer_panorama, eval_baseline, vis_data + cli.add_command(app.main, name='app') + cli.add_command(infer.main, name='infer') + cli.add_command(infer_baseline.main, name='infer_baseline') + cli.add_command(infer_panorama.main, name='infer_panorama') + cli.add_command(eval_baseline.main, name='eval_baseline') + cli.add_command(vis_data.main, name='vis_data') + cli() + + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/moge/scripts/eval_baseline.py b/moge/scripts/eval_baseline.py new file mode 100644 index 0000000..8217d9e --- /dev/null +++ b/moge/scripts/eval_baseline.py @@ -0,0 +1,165 @@ +import os +import sys +from pathlib import Path +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) +import json +from typing import * +import importlib +import importlib.util + +import click + + +@click.command(context_settings={"allow_extra_args": True, "ignore_unknown_options": True}, help='Evaluation script.') +@click.option('--baseline', 'baseline_code_path', type=click.Path(), required=True, help='Path to the baseline model python code.') +@click.option('--config', 'config_path', type=click.Path(), default='configs/eval/all_benchmarks.json', help='Path to the evaluation configurations. ' + 'Defaults to "configs/eval/all_benchmarks.json".') +@click.option('--output', '-o', 'output_path', type=click.Path(), required=True, help='Path to the output json file.') +@click.option('--oracle', 'oracle_mode', is_flag=True, help='Use oracle mode for evaluation, i.e., use the GT intrinsics input.') +@click.option('--dump_pred', is_flag=True, help='Dump predition results.') +@click.option('--dump_gt', is_flag=True, help='Dump ground truth.') +@click.pass_context +def main(ctx: click.Context, baseline_code_path: str, config_path: str, oracle_mode: bool, output_path: Union[str, Path], dump_pred: bool, dump_gt: bool): + # Lazy import + import cv2 + import numpy as np + from tqdm import tqdm + import torch + import torch.nn.functional as F + import utils3d + + from moge.test.baseline import MGEBaselineInterface + from moge.test.dataloader import EvalDataLoaderPipeline + from moge.test.metrics import compute_metrics + from moge.utils.geometry_torch import intrinsics_to_fov + from moge.utils.vis import colorize_depth, colorize_normal + from moge.utils.tools import key_average, flatten_nested_dict, timeit, import_file_as_module + + # Load the baseline model + module = import_file_as_module(baseline_code_path, Path(baseline_code_path).stem) + baseline_cls: Type[MGEBaselineInterface] = getattr(module, 'Baseline') + baseline : MGEBaselineInterface = baseline_cls.load.main(ctx.args, standalone_mode=False) + + # Load the evaluation configurations + with open(config_path, 'r') as f: + config = json.load(f) + + Path(output_path).parent.mkdir(parents=True, exist_ok=True) + all_metrics = {} + # Iterate over the dataset + for benchmark_name, benchmark_config in tqdm(list(config.items()), desc='Benchmarks'): + filenames, metrics_list = [], [] + with ( + EvalDataLoaderPipeline(**benchmark_config) as eval_data_pipe, + tqdm(total=len(eval_data_pipe), desc=benchmark_name, leave=False) as pbar + ): + # Iterate over the samples in the dataset + for i in range(len(eval_data_pipe)): + sample = eval_data_pipe.get() + sample = {k: v.to(baseline.device) if isinstance(v, torch.Tensor) else v for k, v in sample.items()} + image = sample['image'] + gt_intrinsics = sample['intrinsics'] + + # Inference + torch.cuda.synchronize() + with torch.inference_mode(), timeit('_inference_timer', verbose=False) as timer: + if oracle_mode: + pred = baseline.infer_for_evaluation(image, gt_intrinsics) + else: + pred = baseline.infer_for_evaluation(image) + torch.cuda.synchronize() + + # Compute metrics + metrics, misc = compute_metrics(pred, sample, vis=dump_pred or dump_gt) + metrics['inference_time'] = timer.time + metrics_list.append(metrics) + + # Dump results + dump_path = Path(output_path.replace(".json", f"_dump"), f'{benchmark_name}', sample['filename'].replace('.zip', '')) + if dump_pred: + dump_path.joinpath('pred').mkdir(parents=True, exist_ok=True) + cv2.imwrite(str(dump_path / 'pred' / 'image.jpg'), cv2.cvtColor((image.cpu().numpy().transpose(1, 2, 0) * 255).astype(np.uint8), cv2.COLOR_RGB2BGR)) + + with Path(dump_path, 'pred', 'metrics.json').open('w') as f: + json.dump(metrics, f, indent=4) + + if 'pred_points' in misc: + points = misc['pred_points'].cpu().numpy() + cv2.imwrite(str(dump_path / 'pred' / 'points.exr'), cv2.cvtColor(points.astype(np.float32), cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + + if 'pred_depth' in misc: + depth = misc['pred_depth'].cpu().numpy() + if 'mask' in pred: + mask = pred['mask'].cpu().numpy() + depth = np.where(mask, depth, np.inf) + cv2.imwrite(str(dump_path / 'pred' / 'depth.png'), cv2.cvtColor(colorize_depth(depth), cv2.COLOR_RGB2BGR)) + + if 'mask' in pred: + mask = pred['mask'].cpu().numpy() + cv2.imwrite(str(dump_path / 'pred' / 'mask.png'), (mask * 255).astype(np.uint8)) + + if 'normal' in pred: + normal = pred['normal'].cpu().numpy() + cv2.imwrite(str(dump_path / 'pred' / 'normal.png'), cv2.cvtColor(colorize_normal(normal), cv2.COLOR_RGB2BGR)) + + if 'intrinsics' in pred: + intrinsics = pred['intrinsics'] + fov_x, fov_y = intrinsics_to_fov(intrinsics) + with open(dump_path / 'pred' / 'fov.json', 'w') as f: + json.dump({ + 'fov_x': np.rad2deg(fov_x.item()), + 'fov_y': np.rad2deg(fov_y.item()), + 'intrinsics': intrinsics.cpu().numpy().tolist(), + }, f) + + if dump_gt: + dump_path.joinpath('gt').mkdir(parents=True, exist_ok=True) + cv2.imwrite(str(dump_path / 'gt' / 'image.jpg'), cv2.cvtColor((image.cpu().numpy().transpose(1, 2, 0) * 255).astype(np.uint8), cv2.COLOR_RGB2BGR)) + + if 'points' in sample: + points = sample['points'] + cv2.imwrite(str(dump_path / 'gt' / 'points.exr'), cv2.cvtColor(points.cpu().numpy().astype(np.float32), cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + + if 'depth' in sample: + depth = sample['depth'] + mask = sample['depth_mask'] + cv2.imwrite(str(dump_path / 'gt' / 'depth.png'), cv2.cvtColor(colorize_depth(depth.cpu().numpy(), mask=mask.cpu().numpy()), cv2.COLOR_RGB2BGR)) + + if 'normal' in sample: + normal = sample['normal'] + cv2.imwrite(str(dump_path / 'gt' / 'normal.png'), cv2.cvtColor(colorize_normal(normal.cpu().numpy()), cv2.COLOR_RGB2BGR)) + + if 'depth_mask' in sample: + mask = sample['depth_mask'] + cv2.imwrite(str(dump_path / 'gt' /'mask.png'), (mask.cpu().numpy() * 255).astype(np.uint8)) + + if 'intrinsics' in sample: + intrinsics = sample['intrinsics'] + fov_x, fov_y = intrinsics_to_fov(intrinsics) + with open(dump_path / 'gt' / 'info.json', 'w') as f: + json.dump({ + 'fov_x': np.rad2deg(fov_x.item()), + 'fov_y': np.rad2deg(fov_y.item()), + 'intrinsics': intrinsics.cpu().numpy().tolist(), + }, f) + + # Save intermediate results + if i % 100 == 0 or i == len(eval_data_pipe) - 1: + Path(output_path).write_text( + json.dumps({ + **all_metrics, + benchmark_name: key_average(metrics_list) + }, indent=4) + ) + pbar.update(1) + + all_metrics[benchmark_name] = key_average(metrics_list) + + # Save final results + all_metrics['mean'] = key_average(list(all_metrics.values())) + Path(output_path).write_text(json.dumps(all_metrics, indent=4)) + + +if __name__ == '__main__': + main() diff --git a/moge/scripts/infer.py b/moge/scripts/infer.py new file mode 100644 index 0000000..09990f3 --- /dev/null +++ b/moge/scripts/infer.py @@ -0,0 +1,170 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +from pathlib import Path +import sys +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) +from typing import * +import itertools +import json +import warnings + +import click + + +@click.command(help='Inference script') +@click.option('--input', '-i', 'input_path', type=click.Path(exists=True), help='Input image or folder path. "jpg" and "png" are supported.') +@click.option('--fov_x', 'fov_x_', type=float, default=None, help='If camera parameters are known, set the horizontal field of view in degrees. Otherwise, MoGe will estimate it.') +@click.option('--output', '-o', 'output_path', default='./output', type=click.Path(), help='Output folder path') +@click.option('--pretrained', 'pretrained_model_name_or_path', type=str, default=None, help='Pretrained model name or path. If not provided, the corresponding default model will be chosen.') +@click.option('--version', 'model_version', type=click.Choice(['v1', 'v2']), default='v2', help='Model version. Defaults to "v2"') +@click.option('--device', 'device_name', type=str, default='cuda', help='Device name (e.g. "cuda", "cuda:0", "cpu"). Defaults to "cuda"') +@click.option('--fp16', 'use_fp16', is_flag=True, help='Use fp16 precision for much faster inference.') +@click.option('--resize', 'resize_to', type=int, default=None, help='Resize the image(s) & output maps to a specific size. Defaults to None (no resizing).') +@click.option('--resolution_level', type=int, default=9, help='An integer [0-9] for the resolution level for inference. \ +Higher value means more tokens and the finer details will be captured, but inference can be slower. \ +Defaults to 9. Note that it is irrelevant to the output size, which is always the same as the input size. \ +`resolution_level` actually controls `num_tokens`. See `num_tokens` for more details.') +@click.option('--num_tokens', type=int, default=None, help='number of tokens used for inference. A integer in the (suggested) range of `[1200, 2500]`. \ +`resolution_level` will be ignored if `num_tokens` is provided. Default: None') +@click.option('--threshold', type=float, default=0.04, help='Threshold for removing edges. Defaults to 0.01. Smaller value removes more edges. "inf" means no thresholding.') +@click.option('--maps', 'save_maps_', is_flag=True, help='Whether to save the output maps (image, point map, depth map, normal map, mask) and fov.') +@click.option('--glb', 'save_glb_', is_flag=True, help='Whether to save the output as a.glb file. The color will be saved as a texture.') +@click.option('--ply', 'save_ply_', is_flag=True, help='Whether to save the output as a.ply file. The color will be saved as vertex colors.') +@click.option('--show', 'show', is_flag=True, help='Whether show the output in a window. Note that this requires pyglet<2 installed as required by trimesh.') +def main( + input_path: str, + fov_x_: float, + output_path: str, + pretrained_model_name_or_path: str, + model_version: str, + device_name: str, + use_fp16: bool, + resize_to: int, + resolution_level: int, + num_tokens: int, + threshold: float, + save_maps_: bool, + save_glb_: bool, + save_ply_: bool, + show: bool, +): + import cv2 + import numpy as np + import torch + from PIL import Image + from tqdm import tqdm + import click + + from moge.model import import_model_class_by_version + from moge.utils.io import save_glb, save_ply + from moge.utils.vis import colorize_depth, colorize_normal + from moge.utils.geometry_numpy import depth_occlusion_edge_numpy + import utils3d + + device = torch.device(device_name) + + include_suffices = ['jpg', 'png', 'jpeg', 'JPG', 'PNG', 'JPEG'] + if Path(input_path).is_dir(): + image_paths = sorted(itertools.chain(*(Path(input_path).rglob(f'*.{suffix}') for suffix in include_suffices))) + else: + image_paths = [Path(input_path)] + + if len(image_paths) == 0: + raise FileNotFoundError(f'No image files found in {input_path}') + + if pretrained_model_name_or_path is None: + DEFAULT_PRETRAINED_MODEL_FOR_EACH_VERSION = { + "v1": "Ruicheng/moge-vitl", + "v2": "Ruicheng/moge-2-vitl-normal", + } + pretrained_model_name_or_path = DEFAULT_PRETRAINED_MODEL_FOR_EACH_VERSION[model_version] + model = import_model_class_by_version(model_version).from_pretrained(pretrained_model_name_or_path).to(device).eval() + if use_fp16: + model.half() + + if not any([save_maps_, save_glb_, save_ply_]): + warnings.warn('No output format specified. Defaults to saving all. Please use "--maps", "--glb", or "--ply" to specify the output.') + save_maps_ = save_glb_ = save_ply_ = True + + for image_path in (pbar := tqdm(image_paths, desc='Inference', disable=len(image_paths) <= 1)): + if not image_path.exists(): + raise FileNotFoundError(f'File {image_path} does not exist.') + image = cv2.cvtColor(cv2.imread(str(image_path)), cv2.COLOR_BGR2RGB) + height, width = image.shape[:2] + if resize_to is not None: + height, width = min(resize_to, int(resize_to * height / width)), min(resize_to, int(resize_to * width / height)) + image = cv2.resize(image, (width, height), cv2.INTER_AREA) + image_tensor = torch.tensor(image / 255, dtype=torch.float32, device=device).permute(2, 0, 1) + + # Inference + output = model.infer(image_tensor, fov_x=fov_x_, resolution_level=resolution_level, num_tokens=num_tokens, use_fp16=use_fp16) + points, depth, mask, intrinsics = output['points'].cpu().numpy(), output['depth'].cpu().numpy(), output['mask'].cpu().numpy(), output['intrinsics'].cpu().numpy() + normal = output['normal'].cpu().numpy() if 'normal' in output else None + + save_path = Path(output_path, image_path.relative_to(input_path).parent, image_path.stem) + save_path.mkdir(exist_ok=True, parents=True) + + # Save images / maps + if save_maps_: + cv2.imwrite(str(save_path / 'image.jpg'), cv2.cvtColor(image, cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(save_path / 'depth_vis.png'), cv2.cvtColor(colorize_depth(depth), cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(save_path / 'depth.exr'), depth, [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + cv2.imwrite(str(save_path / 'mask.png'), (mask * 255).astype(np.uint8)) + cv2.imwrite(str(save_path / 'points.exr'), cv2.cvtColor(points, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + if normal is not None: + cv2.imwrite(str(save_path / 'normal.png'), cv2.cvtColor(colorize_normal(normal), cv2.COLOR_RGB2BGR)) + fov_x, fov_y = utils3d.np.intrinsics_to_fov(intrinsics) + with open(save_path / 'fov.json', 'w') as f: + json.dump({ + 'fov_x': round(float(np.rad2deg(fov_x)), 2), + 'fov_y': round(float(np.rad2deg(fov_y)), 2), + }, f) + + # Export mesh & visulization + if save_glb_ or save_ply_ or show: + mask_cleaned = mask & ~utils3d.np.depth_map_edge(depth, rtol=threshold) + if normal is None: + faces, vertices, vertex_colors, vertex_uvs = utils3d.np.build_mesh_from_map( + points, + image.astype(np.float32) / 255, + utils3d.np.uv_map(height, width), + mask=mask_cleaned, + tri=True + ) + vertex_normals = None + else: + faces, vertices, vertex_colors, vertex_uvs, vertex_normals = utils3d.np.build_mesh_from_map( + points, + image.astype(np.float32) / 255, + utils3d.np.uv_map(height, width), + normal, + mask=mask_cleaned, + tri=True + ) + # When exporting the model, follow the OpenGL coordinate conventions: + # - world coordinate system: x right, y up, z backward. + # - texture coordinate system: (0, 0) for left-bottom, (1, 1) for right-top. + vertices, vertex_uvs = vertices * [1, -1, -1], vertex_uvs * [1, -1] + [0, 1] + if normal is not None: + vertex_normals = vertex_normals * [1, -1, -1] + + if save_glb_: + save_glb(save_path / 'mesh.glb', vertices, faces, vertex_uvs, image, vertex_normals) + + if save_ply_: + save_ply(save_path / 'pointcloud.ply', vertices, np.zeros((0, 3), dtype=np.int32), vertex_colors, vertex_normals) + + if show: + import trimesh + trimesh.Trimesh( + vertices=vertices, + vertex_colors=vertex_colors, + vertex_normals=vertex_normals, + faces=faces, + process=False + ).show() + + +if __name__ == '__main__': + main() diff --git a/moge/scripts/infer_baseline.py b/moge/scripts/infer_baseline.py new file mode 100644 index 0000000..ef81bc4 --- /dev/null +++ b/moge/scripts/infer_baseline.py @@ -0,0 +1,140 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +from pathlib import Path +import sys +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) +import json +from pathlib import Path +from typing import * +import itertools +import warnings + +import click + + +@click.command(context_settings={"allow_extra_args": True, "ignore_unknown_options": True}, help='Inference script for wrapped baselines methods') +@click.option('--baseline', 'baseline_code_path', required=True, type=click.Path(), help='Path to the baseline model python code.') +@click.option('--input', '-i', 'input_path', type=str, required=True, help='Input image or folder') +@click.option('--output', '-o', 'output_path', type=str, default='./output', help='Output folder') +@click.option('--size', 'image_size', type=int, default=None, help='Resize input image') +@click.option('--skip', is_flag=True, help='Skip existing output') +@click.option('--maps', 'save_maps_', is_flag=True, help='Save output point / depth maps') +@click.option('--ply', 'save_ply_', is_flag=True, help='Save mesh in PLY format') +@click.option('--glb', 'save_glb_', is_flag=True, help='Save mesh in GLB format') +@click.option('--threshold', type=float, default=0.03, help='Depth edge detection threshold for saving mesh') +@click.pass_context +def main(ctx: click.Context, baseline_code_path: str, input_path: str, output_path: str, image_size: int, skip: bool, save_maps_, save_ply_: bool, save_glb_: bool, threshold: float): + # Lazy import + import cv2 + import numpy as np + from tqdm import tqdm + import torch + import utils3d + + from moge.utils.io import save_ply, save_glb + from moge.utils.geometry_numpy import intrinsics_to_fov_numpy + from moge.utils.vis import colorize_depth, colorize_depth_affine, colorize_disparity + from moge.utils.tools import key_average, flatten_nested_dict, timeit, import_file_as_module + from moge.test.baseline import MGEBaselineInterface + + # Load the baseline model + module = import_file_as_module(baseline_code_path, Path(baseline_code_path).stem) + baseline_cls: Type[MGEBaselineInterface] = getattr(module, 'Baseline') + baseline : MGEBaselineInterface = baseline_cls.load.main(ctx.args, standalone_mode=False) + + # Input images list + include_suffices = ['jpg', 'png', 'jpeg', 'JPG', 'PNG', 'JPEG'] + if Path(input_path).is_dir(): + image_paths = sorted(itertools.chain(*(Path(input_path).rglob(f'*.{suffix}') for suffix in include_suffices))) + else: + image_paths = [Path(input_path)] + + if not any([save_maps_, save_glb_, save_ply_]): + warnings.warn('No output format specified. Defaults to saving maps only. Please use "--maps", "--glb", or "--ply" to specify the output.') + save_maps_ = True + + for image_path in (pbar := tqdm(image_paths, desc='Inference', disable=len(image_paths) <= 1)): + # Load one image at a time + image_np = cv2.cvtColor(cv2.imread(str(image_path)), cv2.COLOR_BGR2RGB) + height, width = image_np.shape[:2] + if image_size is not None and max(image_np.shape[:2]) > image_size: + height, width = min(image_size, int(image_size * height / width)), min(image_size, int(image_size * width / height)) + image_np = cv2.resize(image_np, (width, height), cv2.INTER_AREA) + image = torch.from_numpy(image_np.astype(np.float32) / 255.0).permute(2, 0, 1).to(baseline.device) + + # Inference + torch.cuda.synchronize() + with torch.inference_mode(), (timer := timeit('Inference', verbose=False, average=True)): + output = baseline.infer(image) + torch.cuda.synchronize() + + inference_time = timer.average_time + pbar.set_postfix({'average inference time': f'{inference_time:.3f}s'}) + + # Save the output + save_path = Path(output_path, image_path.relative_to(input_path).parent, image_path.stem) + if skip and save_path.exists(): + continue + save_path.mkdir(parents=True, exist_ok=True) + + if save_maps_: + cv2.imwrite(str(save_path / 'image.jpg'), cv2.cvtColor(image_np, cv2.COLOR_RGB2BGR)) + + if 'mask' in output: + mask = output['mask'].cpu().numpy() + cv2.imwrite(str(save_path /'mask.png'), (mask * 255).astype(np.uint8)) + + for k in ['points_metric', 'points_scale_invariant', 'points_affine_invariant']: + if k in output: + points = output[k].cpu().numpy() + cv2.imwrite(str(save_path / f'{k}.exr'), cv2.cvtColor(points, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + + for k in ['depth_metric', 'depth_scale_invariant', 'depth_affine_invariant', 'disparity_affine_invariant']: + if k in output: + depth = output[k].cpu().numpy() + cv2.imwrite(str(save_path / f'{k}.exr'), depth, [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + if k in ['depth_metric', 'depth_scale_invariant']: + depth_vis = colorize_depth(depth) + elif k == 'depth_affine_invariant': + depth_vis = colorize_depth_affine(depth) + elif k == 'disparity_affine_invariant': + depth_vis = colorize_disparity(depth) + cv2.imwrite(str(save_path / f'{k}_vis.png'), cv2.cvtColor(depth_vis, cv2.COLOR_RGB2BGR)) + + if 'intrinsics' in output: + intrinsics = output['intrinsics'].cpu().numpy() + fov_x, fov_y = intrinsics_to_fov_numpy(intrinsics) + with open(save_path / 'fov.json', 'w') as f: + json.dump({ + 'fov_x': float(np.rad2deg(fov_x)), + 'fov_y': float(np.rad2deg(fov_y)), + 'intrinsics': intrinsics.tolist() + }, f, indent=4) + + # Export mesh & visulization + if save_ply_ or save_glb_: + assert any(k in output for k in ['points_metric', 'points_scale_invariant', 'points_affine_invariant']), 'No point map found in output' + points = next(output[k] for k in ['points_metric', 'points_scale_invariant', 'points_affine_invariant'] if k in output).cpu().numpy() + mask = output['mask'] if 'mask' in output else np.ones_like(points[..., 0], dtype=bool) + normals, normals_mask = utils3d.np.point_map_to_normal_map(points, mask=mask) + faces, vertices, vertex_colors, vertex_uvs = utils3d.np.build_mesh_from_map( + points, + image_np.astype(np.float32) / 255, + utils3d.np.uv_map(height, width), + mask=mask & ~(utils3d.np.depth_map_edge(depth, rtol=threshold, mask=mask) & utils3d.np.normal_map_edge(normals, tol=5, mask=normals_mask)), + tri=True + ) + # When exporting the model, follow the OpenGL coordinate conventions: + # - world coordinate system: x right, y up, z backward. + # - texture coordinate system: (0, 0) for left-bottom, (1, 1) for right-top. + vertices, vertex_uvs = vertices * [1, -1, -1], vertex_uvs * [1, -1] + [0, 1] + + if save_glb_: + save_glb(save_path / 'mesh.glb', vertices, faces, vertex_uvs, image_np) + + if save_ply_: + save_ply(save_path / 'mesh.ply', vertices, faces, vertex_colors) + +if __name__ == '__main__': + main() diff --git a/moge/scripts/infer_panorama.py b/moge/scripts/infer_panorama.py new file mode 100644 index 0000000..525a8ad --- /dev/null +++ b/moge/scripts/infer_panorama.py @@ -0,0 +1,162 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +from pathlib import Path +import sys +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) +from typing import * +import itertools +import json +import warnings + +import click + + +@click.command(help='Inference script for panorama images') +@click.option('--input', '-i', 'input_path', type=click.Path(exists=True), required=True, help='Input image or folder path. "jpg" and "png" are supported.') +@click.option('--output', '-o', 'output_path', type=click.Path(), default='./output', help='Output folder path') +@click.option('--pretrained', 'pretrained_model_name_or_path', type=str, default='Ruicheng/moge-vitl', help='Pretrained model name or path. Defaults to "Ruicheng/moge-vitl"') +@click.option('--device', 'device_name', type=str, default='cuda', help='Device name (e.g. "cuda", "cuda:0", "cpu"). Defaults to "cuda"') +@click.option('--resize', 'resize_to', type=int, default=None, help='Resize the image(s) & output maps to a specific size. Defaults to None (no resizing).') +@click.option('--resolution_level', type=int, default=9, help='An integer [0-9] for the resolution level of inference. The higher, the better but slower. Defaults to 9. Note that it is irrelevant to the output resolution.') +@click.option('--threshold', type=float, default=0.03, help='Threshold for removing edges. Defaults to 0.03. Smaller value removes more edges. "inf" means no thresholding.') +@click.option('--batch_size', type=int, default=4, help='Batch size for inference. Defaults to 4.') +@click.option('--splitted', 'save_splitted', is_flag=True, help='Whether to save the splitted images. Defaults to False.') +@click.option('--maps', 'save_maps_', is_flag=True, help='Whether to save the output maps and fov(image, depth, mask, points, fov).') +@click.option('--glb', 'save_glb_', is_flag=True, help='Whether to save the output as a.glb file. The color will be saved as a texture.') +@click.option('--ply', 'save_ply_', is_flag=True, help='Whether to save the output as a.ply file. The color will be saved as vertex colors.') +@click.option('--show', 'show', is_flag=True, help='Whether show the output in a window. Note that this requires pyglet<2 installed as required by trimesh.') +def main( + input_path: str, + output_path: str, + pretrained_model_name_or_path: str, + device_name: str, + resize_to: int, + resolution_level: int, + threshold: float, + batch_size: int, + save_splitted: bool, + save_maps_: bool, + save_glb_: bool, + save_ply_: bool, + show: bool, +): + # Lazy import + import cv2 + import numpy as np + from numpy import ndarray + import torch + from PIL import Image + from tqdm import tqdm, trange + import trimesh + import trimesh.visual + from scipy.sparse import csr_array, hstack, vstack + from scipy.ndimage import convolve + from scipy.sparse.linalg import lsmr + + import utils3d + from moge.model.v1 import MoGeModel + from moge.utils.io import save_glb, save_ply + from moge.utils.vis import colorize_depth + from moge.utils.panorama import spherical_uv_to_directions, get_panorama_cameras, split_panorama_image, merge_panorama_depth + + + device = torch.device(device_name) + + include_suffices = ['jpg', 'png', 'jpeg', 'JPG', 'PNG', 'JPEG'] + if Path(input_path).is_dir(): + image_paths = sorted(itertools.chain(*(Path(input_path).rglob(f'*.{suffix}') for suffix in include_suffices))) + else: + image_paths = [Path(input_path)] + + if len(image_paths) == 0: + raise FileNotFoundError(f'No image files found in {input_path}') + + # Write outputs + if not any([save_maps_, save_glb_, save_ply_]): + warnings.warn('No output format specified. Defaults to saving all. Please use "--maps", "--glb", or "--ply" to specify the output.') + save_maps_ = save_glb_ = save_ply_ = True + + model = MoGeModel.from_pretrained(pretrained_model_name_or_path).to(device).eval() + + for image_path in (pbar := tqdm(image_paths, desc='Total images', disable=len(image_paths) <= 1)): + image = cv2.cvtColor(cv2.imread(str(image_path)), cv2.COLOR_BGR2RGB) + height, width = image.shape[:2] + if resize_to is not None: + height, width = min(resize_to, int(resize_to * height / width)), min(resize_to, int(resize_to * width / height)) + image = cv2.resize(image, (width, height), cv2.INTER_AREA) + + splitted_extrinsics, splitted_intriniscs = get_panorama_cameras() + splitted_resolution = 512 + splitted_images = split_panorama_image(image, splitted_extrinsics, splitted_intriniscs, splitted_resolution) + + # Infer each view + print('Inferring...') if pbar.disable else pbar.set_postfix_str(f'Inferring') + + splitted_distance_maps, splitted_masks = [], [] + for i in trange(0, len(splitted_images), batch_size, desc='Inferring splitted views', disable=len(splitted_images) <= batch_size, leave=False): + image_tensor = torch.tensor(np.stack(splitted_images[i:i + batch_size]) / 255, dtype=torch.float32, device=device).permute(0, 3, 1, 2) + fov_x, fov_y = np.rad2deg(utils3d.np.intrinsics_to_fov(np.array(splitted_intriniscs[i:i + batch_size]))) + fov_x = torch.tensor(fov_x, dtype=torch.float32, device=device) + output = model.infer(image_tensor, fov_x=fov_x, apply_mask=False) + distance_map, mask = output['points'].norm(dim=-1).cpu().numpy(), output['mask'].cpu().numpy() + splitted_distance_maps.extend(list(distance_map)) + splitted_masks.extend(list(mask)) + + # Save splitted + if save_splitted: + splitted_save_path = Path(output_path, image_path.stem, 'splitted') + splitted_save_path.mkdir(exist_ok=True, parents=True) + for i in range(len(splitted_images)): + cv2.imwrite(str(splitted_save_path / f'{i:02d}.jpg'), cv2.cvtColor(splitted_images[i], cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(splitted_save_path / f'{i:02d}_distance_vis.png'), cv2.cvtColor(colorize_depth(splitted_distance_maps[i], splitted_masks[i]), cv2.COLOR_RGB2BGR)) + + # Merge + print('Merging...') if pbar.disable else pbar.set_postfix_str(f'Merging') + + merging_width, merging_height = min(1920, width), min(960, height) + panorama_depth, panorama_mask = merge_panorama_depth(merging_width, merging_height, splitted_distance_maps, splitted_masks, splitted_extrinsics, splitted_intriniscs) + panorama_depth = panorama_depth.astype(np.float32) + panorama_depth = cv2.resize(panorama_depth, (width, height), cv2.INTER_LINEAR) + panorama_mask = cv2.resize(panorama_mask.astype(np.uint8), (width, height), cv2.INTER_NEAREST) > 0 + points = panorama_depth[:, :, None] * spherical_uv_to_directions(utils3d.np.uv_map(height, width)) + + # Write outputs + print('Writing outputs...') if pbar.disable else pbar.set_postfix_str(f'Inferring') + save_path = Path(output_path, image_path.relative_to(input_path).parent, image_path.stem) + save_path.mkdir(exist_ok=True, parents=True) + if save_maps_: + cv2.imwrite(str(save_path / 'image.jpg'), cv2.cvtColor(image, cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(save_path / 'depth_vis.png'), cv2.cvtColor(colorize_depth(panorama_depth, mask=panorama_mask), cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(save_path / 'depth.exr'), panorama_depth, [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + cv2.imwrite(str(save_path / 'points.exr'), points, [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + cv2.imwrite(str(save_path /'mask.png'), (panorama_mask * 255).astype(np.uint8)) + + # Export mesh & visulization + if save_glb_ or save_ply_ or show: + normals, normals_mask = utils3d.np.point_map_to_normal_map(points, panorama_mask) + faces, vertices, vertex_colors, vertex_uvs = utils3d.np.build_mesh_from_map( + points, + image.astype(np.float32) / 255, + utils3d.np.uv_map(height, width), + mask=panorama_mask & ~(utils3d.np.depth_map_edge(panorama_depth, rtol=threshold) & utils3d.np.normal_map_edge(normals, tol=5, mask=normals_mask)), + tri=True + ) + + if save_glb_: + save_glb(save_path / 'mesh.glb', vertices, faces, vertex_uvs, image) + + if save_ply_: + save_ply(save_path / 'mesh.ply', vertices, faces, vertex_colors) + + if show: + trimesh.Trimesh( + vertices=vertices, + vertex_colors=vertex_colors, + faces=faces, + process=False + ).show() + + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/moge/scripts/train.py b/moge/scripts/train.py new file mode 100644 index 0000000..6d810cd --- /dev/null +++ b/moge/scripts/train.py @@ -0,0 +1,461 @@ +import os +from pathlib import Path +import sys +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) +import json +import time +import random +from typing import * +import itertools +from contextlib import nullcontext +from concurrent.futures import ThreadPoolExecutor +import io + +import numpy as np +import cv2 +from PIL import Image +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.version +import accelerate +from accelerate import Accelerator, DistributedDataParallelKwargs +from accelerate.utils import set_seed +import utils3d +import click +from tqdm import tqdm, trange +import mlflow +torch.backends.cudnn.benchmark = False # Varying input size, make sure cudnn benchmark is disabled + +from moge.train.dataloader import TrainDataLoaderPipeline +from moge.train.losses import ( + affine_invariant_global_loss, + affine_invariant_local_loss, + edge_loss, + normal_loss, + mask_l2_loss, + mask_bce_loss, + metric_scale_loss, + normal_map_loss, + monitoring, +) +from moge.train.utils import build_optimizer, build_lr_scheduler +from moge.utils.geometry_torch import intrinsics_to_fov +from moge.utils.vis import colorize_depth, colorize_normal +from moge.utils.tools import key_average, recursive_replace, CallbackOnException, flatten_nested_dict +from moge.test.metrics import compute_metrics + + +@click.command() +@click.option('--config', 'config_path', type=str, default='configs/debug.json') +@click.option('--workspace', type=str, default='workspace/debug', help='Path to the workspace') +@click.option('--checkpoint', 'checkpoint_path', type=str, default=None, help='Path to the checkpoint to load. "latest" to load latest checkpoint in workspace, integer to load by step number') +@click.option('--batch_size_forward', type=int, default=8, help='Batch size for each forward pass on each device') +@click.option('--gradient_accumulation_steps', type=int, default=1, help='Number of steps to accumulate gradients') +@click.option('--enable_gradient_checkpointing', type=bool, default=True, help='Use gradient checkpointing in backbone') +@click.option('--enable_mixed_precision', type=bool, default=False, help='Use mixed precision training. Backbone is converted to FP16') +@click.option('--enable_ema', type=bool, default=True, help='Maintain an exponential moving average of the model weights') +@click.option('--num_iterations', type=int, default=1000000, help='Number of iterations to train the model') +@click.option('--save_every', type=int, default=10000, help='Save checkpoint every n iterations') +@click.option('--log_every', type=int, default=1000, help='Log metrics every n iterations') +@click.option('--vis_every', type=int, default=0, help='Visualize every n iterations') +@click.option('--num_vis_images', type=int, default=32, help='Number of images to visualize, must be a multiple of divided batch size') +@click.option('--enable_mlflow', type=bool, default=True, help='Log metrics to MLFlow') +@click.option('--seed', type=int, default=0, help='Random seed') +def main( + config_path: str, + workspace: str, + checkpoint_path: str, + batch_size_forward: int, + gradient_accumulation_steps: int, + enable_gradient_checkpointing: bool, + enable_mixed_precision: bool, + enable_ema: bool, + num_iterations: int, + save_every: int, + log_every: int, + vis_every: int, + num_vis_images: int, + enable_mlflow: bool, + seed: Optional[int], +): + # Load config + with open(config_path, 'r') as f: + config = json.load(f) + + accelerator = Accelerator( + gradient_accumulation_steps=gradient_accumulation_steps, + mixed_precision='fp16' if enable_mixed_precision else None, + kwargs_handlers=[ + DistributedDataParallelKwargs(find_unused_parameters=True) + ] + ) + device = accelerator.device + batch_size_total = batch_size_forward * gradient_accumulation_steps * accelerator.num_processes + + # Log config + if accelerator.is_main_process: + if enable_mlflow: + try: + mlflow.log_params({ + **click.get_current_context().params, + 'batch_size_total': batch_size_total, + }) + except: + print('Failed to log config to MLFlow') + Path(workspace).mkdir(parents=True, exist_ok=True) + with Path(workspace).joinpath('config.json').open('w') as f: + json.dump(config, f, indent=4) + + # Set seed + if seed is not None: + set_seed(seed, device_specific=True) + + # Initialize model + print('Initialize model') + with accelerator.local_main_process_first(): + from moge.model import import_model_class_by_version + MoGeModel = import_model_class_by_version(config['model_version']) + model = MoGeModel(**config['model']) + count_total_parameters = sum(p.numel() for p in model.parameters()) + print(f'Total parameters: {count_total_parameters}') + + # Set up EMA model + if enable_ema and accelerator.is_main_process: + ema_avg_fn = lambda averaged_model_parameter, model_parameter, num_averaged: 0.999 * averaged_model_parameter + 0.001 * model_parameter + ema_model = torch.optim.swa_utils.AveragedModel(model, device=accelerator.device, avg_fn=ema_avg_fn) + + # Set gradient checkpointing + if enable_gradient_checkpointing: + model.enable_gradient_checkpointing() + import warnings + warnings.filterwarnings("ignore", category=FutureWarning, module="torch.utils.checkpoint") + + # Initalize optimizer & lr scheduler + optimizer = build_optimizer(model, config['optimizer']) + lr_scheduler = build_lr_scheduler(optimizer, config['lr_scheduler']) + + count_grouped_parameters = [sum(p.numel() for p in param_group['params'] if p.requires_grad) for param_group in optimizer.param_groups] + for i, count in enumerate(count_grouped_parameters): + print(f'- Group {i}: {count} parameters') + + # Attempt to load checkpoint + checkpoint: Dict[str, Any] + with accelerator.local_main_process_first(): + if checkpoint_path is None: + # - No checkpoint + checkpoint = None + elif checkpoint_path.endswith('.pt'): + # - Load specific checkpoint file + print(f'Load checkpoint: {checkpoint_path}') + checkpoint = torch.load(checkpoint_path, map_location='cpu', weights_only=True) + elif checkpoint_path == "latest": + # - Load latest checkpoint + checkpoint_path = Path(workspace, 'checkpoint', 'latest.pt') + if checkpoint_path.exists(): + print(f'Load checkpoint: {checkpoint_path}') + checkpoint = torch.load(checkpoint_path, map_location='cpu', weights_only=True) + i_step = checkpoint['step'] + if 'model' not in checkpoint and (checkpoint_model_path := Path(workspace, 'checkpoint', f'{i_step:08d}.pt')).exists(): + print(f'Load model checkpoint: {checkpoint_model_path}') + checkpoint['model'] = torch.load(checkpoint_model_path, map_location='cpu', weights_only=True)['model'] + if 'optimizer' not in checkpoint and (checkpoint_optimizer_path := Path(workspace, 'checkpoint', f'{i_step:08d}_optimizer.pt')).exists(): + print(f'Load optimizer checkpoint: {checkpoint_optimizer_path}') + checkpoint.update(torch.load(checkpoint_optimizer_path, map_location='cpu', weights_only=True)) + if enable_ema and accelerator.is_main_process: + if 'ema_model' not in checkpoint and (checkpoint_ema_model_path := Path(workspace, 'checkpoint', f'{i_step:08d}_ema.pt')).exists(): + print(f'Load EMA model checkpoint: {checkpoint_ema_model_path}') + checkpoint['ema_model'] = torch.load(checkpoint_ema_model_path, map_location='cpu', weights_only=True)['model'] + else: + print(f'No latest checkpoint found. Start from scratch.') + checkpoint = None + else: + # - Load by step number + i_step = int(checkpoint_path) + checkpoint = {'step': i_step} + if (checkpoint_model_path := Path(workspace, 'checkpoint', f'{i_step:08d}.pt')).exists(): + print(f'Load model checkpoint: {checkpoint_model_path}') + checkpoint['model'] = torch.load(checkpoint_model_path, map_location='cpu', weights_only=True)['model'] + if (checkpoint_optimizer_path := Path(workspace, 'checkpoint', f'{i_step:08d}_optimizer.pt')).exists(): + print(f'Load optimizer checkpoint: {checkpoint_optimizer_path}') + checkpoint.update(torch.load(checkpoint_optimizer_path, map_location='cpu', weights_only=True)) + if enable_ema and accelerator.is_main_process: + if (checkpoint_ema_model_path := Path(workspace, 'checkpoint', f'{i_step:08d}_ema.pt')).exists(): + print(f'Load EMA model checkpoint: {checkpoint_ema_model_path}') + checkpoint['ema_model'] = torch.load(checkpoint_ema_model_path, map_location='cpu', weights_only=True)['model'] + + if checkpoint is None: + # Initialize model weights + print('Initialize model weights') + with accelerator.local_main_process_first(): + model.init_weights() + initial_step = 0 + else: + model.load_state_dict(checkpoint['model'], strict=False) + if 'step' in checkpoint: + initial_step = checkpoint['step'] + 1 + else: + initial_step = 0 + if 'optimizer' in checkpoint: + optimizer.load_state_dict(checkpoint['optimizer']) + if enable_ema and accelerator.is_main_process and 'ema_model' in checkpoint: + ema_model.module.load_state_dict(checkpoint['ema_model'], strict=False) + if 'lr_scheduler' in checkpoint: + lr_scheduler.load_state_dict(checkpoint['lr_scheduler']) + + del checkpoint + + model, optimizer = accelerator.prepare(model, optimizer) + if torch.version.hip and isinstance(model, torch.nn.parallel.DistributedDataParallel): + # Hacking potential gradient synchronization issue in ROCm backend + from moge.model.utils import sync_ddp_hook + model.register_comm_hook(None, sync_ddp_hook) + + # Initialize training data pipeline + with accelerator.local_main_process_first(): + train_data_pipe = TrainDataLoaderPipeline(config['data'], batch_size_forward) + + def _write_bytes_retry_loop(save_path: Path, data: bytes): + while True: + try: + save_path.write_bytes(data) + break + except Exception as e: + print('Error while saving checkpoint, retrying in 1 minute: ', e) + time.sleep(60) + + # Ready to train + records = [] + model.train() + with ( + train_data_pipe, + tqdm(initial=initial_step, total=num_iterations, desc='Training', disable=not accelerator.is_main_process) as pbar, + ThreadPoolExecutor(max_workers=1) as save_checkpoint_executor, + ): + # Get some batches for visualization + if accelerator.is_main_process: + batches_for_vis: List[Dict[str, torch.Tensor]] = [] + num_vis_images = num_vis_images // batch_size_forward * batch_size_forward + for _ in range(num_vis_images // batch_size_forward): + batch = train_data_pipe.get() + batches_for_vis.append(batch) + + # Visualize GT + if vis_every > 0 and accelerator.is_main_process and initial_step == 0: + save_dir = Path(workspace).joinpath('vis/gt') + for i_batch, batch in enumerate(tqdm(batches_for_vis, desc='Visualize GT', leave=False)): + image, gt_depth, gt_normal, gt_intrinsics, info = batch['image'], batch['depth'], batch['normal'], batch['intrinsics'], batch['info'] + gt_points = utils3d.pt.depth_map_to_point_map(gt_depth, intrinsics=gt_intrinsics) + for i_instance in range(batch['image'].shape[0]): + idx = i_batch * batch_size_forward + i_instance + image_i = (image[i_instance].numpy().transpose(1, 2, 0) * 255).astype(np.uint8) + gt_depth_i = gt_depth[i_instance].numpy() + gt_points_i = gt_points[i_instance].numpy() + gt_normal_i = gt_normal[i_instance].numpy() + save_dir.joinpath(f'{idx:04d}').mkdir(parents=True, exist_ok=True) + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/image.jpg')), cv2.cvtColor(image_i, cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/points.exr')), cv2.cvtColor(gt_points_i, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/depth_vis.png')), cv2.cvtColor(colorize_depth(gt_depth_i), cv2.COLOR_RGB2BGR)) + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/normal.png')), cv2.cvtColor(colorize_normal(gt_normal_i), cv2.COLOR_RGB2BGR)) + with save_dir.joinpath(f'{idx:04d}/info.json').open('w') as f: + json.dump(info[i_instance], f) + + # Reset seed to avoid training on the same data when resuming training + if seed is not None: + set_seed(seed + initial_step, device_specific=True) + + # Training loop + for i_step in range(initial_step, num_iterations): + + i_accumulate, weight_accumulate = 0, 0 + while i_accumulate < gradient_accumulation_steps: + # Load batch + batch = train_data_pipe.get() + image, gt_depth, gt_normal, gt_mask_fin, gt_mask_inf, gt_intrinsics, label_type, is_metric = batch['image'], batch['depth'], batch['normal'], batch['depth_mask_fin'], batch['depth_mask_inf'], batch['intrinsics'], batch['label_type'], batch['is_metric'] + image, gt_depth, gt_normal, gt_mask_fin, gt_mask_inf, gt_intrinsics = image.to(device), gt_depth.to(device), gt_normal.to(device), gt_mask_fin.to(device), gt_mask_inf.to(device), gt_intrinsics.to(device) + current_batch_size = image.shape[0] + if all(label == 'invalid' for label in label_type): + continue # NOTE: Skip all-invalid batches to avoid messing up the optimizer. + + gt_points = utils3d.pt.depth_map_to_point_map(gt_depth, intrinsics=gt_intrinsics) + gt_focal = 1 / (1 / gt_intrinsics[..., 0, 0] ** 2 + 1 / gt_intrinsics[..., 1, 1] ** 2) ** 0.5 + + with accelerator.accumulate(model): + # Forward + if i_step <= config.get('low_resolution_training_steps', 0): + num_tokens = config['model']['num_tokens_range'][0] + else: + num_tokens = accelerate.utils.broadcast_object_list([random.randint(*config['model']['num_tokens_range'])])[0] + with torch.autocast(device_type=accelerator.device.type, dtype=torch.float16, enabled=enable_mixed_precision): + output = model(image, num_tokens=num_tokens) + pred_points, pred_mask, pred_normal, pred_metric_scale = (output.get(k, None) for k in ['points', 'mask', 'normal', 'metric_scale']) + + # Compute loss (per instance) + loss_list, weight_list = [], [] + for i in range(current_batch_size): + gt_metric_scale = None + loss_dict, weight_dict, misc_dict = {}, {}, {} + misc_dict['monitoring'] = monitoring(pred_points[i]) + for k, v in config['loss'][label_type[i]].items(): + weight_dict[k] = v['weight'] + if v['function'] == 'affine_invariant_global_loss': + loss_dict[k], misc_dict[k], gt_metric_scale = affine_invariant_global_loss(pred_points[i], gt_points[i], **v['params']) + elif v['function'] == 'affine_invariant_local_loss': + loss_dict[k], misc_dict[k] = affine_invariant_local_loss(pred_points[i], gt_points[i], gt_focal[i], gt_metric_scale, **v['params']) + elif v['function'] == 'normal_loss': + loss_dict[k], misc_dict[k] = normal_loss(pred_points[i], gt_points[i]) + elif v['function'] == 'edge_loss': + loss_dict[k], misc_dict[k] = edge_loss(pred_points[i], gt_points[i]) + elif v['function'] == 'normal_map_loss': + loss_dict[k], misc_dict[k] = normal_map_loss(pred_normal[i], gt_normal[i]) + elif v['function'] == 'mask_bce_loss': + loss_dict[k], misc_dict[k] = mask_bce_loss(pred_mask[i], gt_mask_fin[i], gt_mask_inf[i]) + elif v['function'] == 'mask_l2_loss': + loss_dict[k], misc_dict[k] = mask_l2_loss(pred_mask[i], gt_mask_fin[i], gt_mask_inf[i]) + elif v['function'] == 'metric_scale_loss': + if is_metric[i] and pred_metric_scale is not None: + loss_dict[k], misc_dict[k] = metric_scale_loss(pred_metric_scale[i], gt_metric_scale) + else: + raise ValueError(f'Undefined loss function: {v["function"]}') + weight_dict = {'.'.join(k): v for k, v in flatten_nested_dict(weight_dict).items()} + loss_dict = {'.'.join(k): v for k, v in flatten_nested_dict(loss_dict).items()} + loss_ = sum([weight_dict[k] * loss_dict[k] for k in loss_dict], start=torch.tensor(0.0, device=device)) + loss_list.append(loss_) + + if torch.isnan(loss_).item(): + pbar.write(f'NaN loss in process {accelerator.process_index}') + pbar.write(str(loss_dict)) + + misc_dict = {'.'.join(k): v for k, v in flatten_nested_dict(misc_dict).items()} + records.append({ + **{k: v.item() for k, v in loss_dict.items()}, + **misc_dict, + }) + + loss = sum(loss_list) / len(loss_list) + + # Backward & update + accelerator.backward(loss) + if accelerator.sync_gradients: + if not enable_mixed_precision and any(torch.isnan(p.grad).any() for p in model.parameters() if p.grad is not None): + if accelerator.is_main_process: + pbar.write(f'NaN gradients, skip update') + optimizer.zero_grad() + continue + accelerator.clip_grad_norm_(model.parameters(), 1.0) + + optimizer.step() + optimizer.zero_grad() + + i_accumulate += 1 + + lr_scheduler.step() + + # EMA update + if enable_ema and accelerator.is_main_process and accelerator.sync_gradients: + ema_model.update_parameters(model) + + # Log metrics + if i_step == initial_step or i_step % log_every == 0: + records = [key_average(records)] + records = accelerator.gather_for_metrics(records, use_gather_object=True) + if accelerator.is_main_process: + records = key_average(records) + if enable_mlflow: + try: + mlflow.log_metrics(records, step=i_step) + except Exception as e: + print(f'Error while logging metrics to mlflow: {e}') + records = [] + + # Save model weight checkpoint + if accelerator.is_main_process and (i_step % save_every == 0): + # NOTE: Writing checkpoint is done in a separate thread to avoid blocking the main process + pbar.write(f'Save checkpoint: {i_step:08d}') + Path(workspace, 'checkpoint').mkdir(parents=True, exist_ok=True) + + # Model checkpoint + with io.BytesIO() as f: + torch.save({ + 'model_config': config['model'], + 'model': accelerator.unwrap_model(model).state_dict(), + }, f) + checkpoint_bytes = f.getvalue() + save_checkpoint_executor.submit( + _write_bytes_retry_loop, Path(workspace, 'checkpoint', f'{i_step:08d}.pt'), checkpoint_bytes + ) + + # Optimizer checkpoint + with io.BytesIO() as f: + torch.save({ + 'model_config': config['model'], + 'step': i_step, + 'optimizer': optimizer.state_dict(), + 'lr_scheduler': lr_scheduler.state_dict(), + }, f) + checkpoint_bytes = f.getvalue() + save_checkpoint_executor.submit( + _write_bytes_retry_loop, Path(workspace, 'checkpoint', f'{i_step:08d}_optimizer.pt'), checkpoint_bytes + ) + + # EMA model checkpoint + if enable_ema: + with io.BytesIO() as f: + torch.save({ + 'model_config': config['model'], + 'model': ema_model.module.state_dict(), + }, f) + checkpoint_bytes = f.getvalue() + save_checkpoint_executor.submit( + _write_bytes_retry_loop, Path(workspace, 'checkpoint', f'{i_step:08d}_ema.pt'), checkpoint_bytes + ) + + # Latest checkpoint + with io.BytesIO() as f: + torch.save({ + 'model_config': config['model'], + 'step': i_step, + }, f) + checkpoint_bytes = f.getvalue() + save_checkpoint_executor.submit( + _write_bytes_retry_loop, Path(workspace, 'checkpoint', 'latest.pt'), checkpoint_bytes + ) + + # Visualize + if vis_every > 0 and accelerator.is_main_process and (i_step == initial_step or i_step % vis_every == 0): + unwrapped_model = accelerator.unwrap_model(model) + save_dir = Path(workspace).joinpath(f'vis/step_{i_step:08d}') + save_dir.mkdir(parents=True, exist_ok=True) + with torch.inference_mode(): + for i_batch, batch in enumerate(tqdm(batches_for_vis, desc=f'Visualize: {i_step:08d}', leave=False)): + image, gt_depth, gt_intrinsics = batch['image'], batch['depth'], batch['intrinsics'] + image, gt_depth, gt_intrinsics = image.to(device), gt_depth.to(device), gt_intrinsics.to(device) + + output = unwrapped_model.infer(image) + pred_points = output['points'].cpu().numpy() if 'points' in output else None + pred_depth = output['depth'].cpu().numpy() if 'depth' in output else None + pred_mask = output['mask'].cpu().numpy() if 'mask' in output else None + pred_normal = output['normal'].cpu().numpy() if 'normal' in output else None + pred_uncertainty = output['uncertainty'].cpu().numpy() if 'uncertainty' in output else None + image = (image.cpu().numpy().transpose(0, 2, 3, 1) * 255).astype(np.uint8) + + for i_instance in range(image.shape[0]): + idx = i_batch * batch_size_forward + i_instance + save_dir.joinpath(f'{idx:04d}').mkdir(parents=True, exist_ok=True) + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/image.jpg')), cv2.cvtColor(image[i_instance], cv2.COLOR_RGB2BGR)) + if pred_points is not None: + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/points.exr')), cv2.cvtColor(pred_points[i_instance], cv2.COLOR_RGB2BGR), [cv2.IMWRITE_EXR_TYPE, cv2.IMWRITE_EXR_TYPE_FLOAT]) + if pred_mask is not None: + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/mask.png')), pred_mask[i_instance] * 255) + if pred_depth is not None: + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/depth_vis.png')), cv2.cvtColor(colorize_depth(pred_depth[i_instance], pred_mask[i_instance] if pred_mask is not None else None), cv2.COLOR_RGB2BGR)) + if pred_normal is not None: + cv2.imwrite(str(save_dir.joinpath(f'{idx:04d}/normal_vis.png')), cv2.cvtColor(colorize_normal(pred_normal[i_instance], pred_mask[i_instance] if pred_mask is not None else None), cv2.COLOR_RGB2BGR)) + + pbar.set_postfix({'loss': loss.item()}, refresh=False) + pbar.update(1) + + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/moge/scripts/vis_data.py b/moge/scripts/vis_data.py new file mode 100644 index 0000000..fcca724 --- /dev/null +++ b/moge/scripts/vis_data.py @@ -0,0 +1,84 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +import sys +from pathlib import Path +if (_package_root := str(Path(__file__).absolute().parents[2])) not in sys.path: + sys.path.insert(0, _package_root) + +import click + + +@click.command() +@click.argument('folder_or_path', type=click.Path(exists=True)) +@click.option('--output', '-o', 'output_folder', type=click.Path(), help='Path to output folder') +@click.option('--max_depth', '-m', type=float, default=float('inf'), help='max depth') +@click.option('--fov', type=float, default=None, help='field of view in degrees') +@click.option('--show', 'show', is_flag=True, help='show point cloud') +@click.option('--depth', 'depth_filename', type=str, default='depth.png', help='depth image file name') +@click.option('--ply', 'save_ply', is_flag=True, help='save point cloud as PLY file') +@click.option('--depth_vis', 'save_depth_vis', is_flag=True, help='save depth image') +@click.option('--inf', 'inf_mask', is_flag=True, help='use infinity mask') +@click.option('--version', 'version', type=str, default='v3', help='version of rgbd data') +def main( + folder_or_path: str, + output_folder: str, + max_depth: float, + fov: float, + depth_filename: str, + show: bool, + save_ply: bool, + save_depth_vis: bool, + inf_mask: bool, + version: str +): + # Lazy import + import cv2 + import numpy as np + import utils3d + from tqdm import tqdm + import trimesh + + from moge.utils.io import read_image, read_depth, read_json + from moge.utils.vis import colorize_depth, colorize_normal + + filepaths = sorted(p.parent for p in Path(folder_or_path).rglob('meta.json')) + + for filepath in tqdm(filepaths): + image = read_image(Path(filepath, 'image.jpg')) + depth = read_depth(Path(filepath, depth_filename)) + meta = read_json(Path(filepath,'meta.json')) + depth_mask = np.isfinite(depth) + depth_mask_inf = (depth == np.inf) + intrinsics = np.array(meta['intrinsics']) + + extrinsics = np.array([[1, 0, 0, 0], [0, -1, 0, 0], [0, 0, -1, 0], [0, 0, 0, 1]], dtype=float) # OpenGL's identity camera + verts = utils3d.np.unproject_cv(utils3d.np.uv_map(image.shape[:2]), depth, extrinsics=extrinsics, intrinsics=intrinsics) + + depth_mask_ply = depth_mask & (depth < depth[depth_mask].min() * max_depth) + point_cloud = trimesh.PointCloud(verts[depth_mask_ply], image[depth_mask_ply] / 255) + + if show: + point_cloud.show() + + if output_folder is None: + output_path = filepath + else: + output_path = Path(output_folder, filepath.name) + output_path.mkdir(exist_ok=True, parents=True) + + if inf_mask: + depth = np.where(depth_mask_inf, np.inf, depth) + depth_mask = depth_mask | depth_mask_inf + + if save_depth_vis: + p = output_path.joinpath('depth_vis.png') + cv2.imwrite(str(p), cv2.cvtColor(colorize_depth(depth, depth_mask), cv2.COLOR_RGB2BGR)) + print(f"{p}") + + if save_ply: + p = output_path.joinpath('pointcloud.ply') + point_cloud.export(p) + print(f"{p}") + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/moge/test/__init__.py b/moge/test/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/moge/test/baseline.py b/moge/test/baseline.py new file mode 100644 index 0000000..05980aa --- /dev/null +++ b/moge/test/baseline.py @@ -0,0 +1,43 @@ +from typing import * + +import click +import torch + + +class MGEBaselineInterface: + """ + Abstract class for model wrapper to uniformize the interface of loading and inference across different models. + """ + device: torch.device + + @click.command() + @staticmethod + def load(*args, **kwargs) -> "MGEBaselineInterface": + """ + Customized static method to create an instance of the model wrapper from command line arguments. Decorated by `click.command()` + """ + raise NotImplementedError(f"{type(self).__name__} has not implemented the load method.") + + def infer(self, image: torch.FloatTensor, intrinsics: Optional[torch.Tensor] = None) -> Dict[str, torch.Tensor]: + """ + ### Parameters + `image`: [B, 3, H, W] or [3, H, W], RGB values in range [0, 1] + `intrinsics`: [B, 3, 3] or [3, 3], camera intrinsics. Optional. + + ### Returns + A dictionary containing: + - `points_*`. point map output in OpenCV identity camera space. + Supported suffixes: `metric`, `scale_invariant`, `affine_invariant`. + - `depth_*`. depth map output + Supported suffixes: `metric` (in meters), `scale_invariant`, `affine_invariant`. + - `disparity_affine_invariant`. affine disparity map output + """ + raise NotImplementedError(f"{type(self).__name__} has not implemented the infer method.") + + def infer_for_evaluation(self, image: torch.FloatTensor, intrinsics: Optional[torch.Tensor] = None) -> Dict[str, torch.Tensor]: + """ + If the model has a special evaluation mode, override this method to provide the evaluation mode inference. + + By default, this method simply calls `infer()`. + """ + return self.infer(image, intrinsics) \ No newline at end of file diff --git a/moge/test/dataloader.py b/moge/test/dataloader.py new file mode 100644 index 0000000..97a9298 --- /dev/null +++ b/moge/test/dataloader.py @@ -0,0 +1,221 @@ +import os +from typing import * +from pathlib import Path +import math + +import numpy as np +import torch +from PIL import Image +import cv2 +import utils3d +import pipeline + +from ..utils.geometry_numpy import focal_to_fov_numpy, norm3d +from ..utils.io import * +from ..utils.tools import timeit + + +class EvalDataLoaderPipeline: + + def __init__( + self, + path: str, + width: int, + height: int, + split: int = '.index.txt', + drop_max_depth: float = 1000., + num_load_workers: int = 4, + num_process_workers: int = 8, + include_segmentation: bool = False, + include_normal: bool = False, + depth_to_normal: bool = False, + max_segments: int = 100, + min_seg_area: int = 1000, + depth_unit: str = None, + has_sharp_boundary = False, + subset: int = None, + ): + filenames = Path(path).joinpath(split).read_text(encoding='utf-8').splitlines() + filenames = filenames[::subset] + self.width = width + self.height = height + self.drop_max_depth = drop_max_depth + self.path = Path(path) + self.filenames = filenames + self.include_segmentation = include_segmentation + self.include_normal = include_normal + self.max_segments = max_segments + self.min_seg_area = min_seg_area + self.depth_to_normal = depth_to_normal + self.depth_unit = depth_unit + self.has_sharp_boundary = has_sharp_boundary + + self.rng = np.random.default_rng(seed=0) + + self.pipeline = pipeline.Sequential([ + self._generator, + pipeline.Parallel([self._load_instance] * num_load_workers), + pipeline.Parallel([self._process_instance] * num_process_workers), + pipeline.Buffer(4) + ]) + + def __len__(self): + return math.ceil(len(self.filenames)) + + def _generator(self): + for idx in range(len(self)): + yield idx + + def _load_instance(self, idx): + if idx >= len(self.filenames): + return None + + path = self.path.joinpath(self.filenames[idx]) + + instance = { + 'filename': self.filenames[idx], + 'width': self.width, + 'height': self.height, + } + instance['image'] = read_image(Path(path, 'image.jpg')) + + depth = read_depth(Path(path, 'depth.png')) # ignore depth unit from depth file, use config instead + instance.update({ + 'depth': np.nan_to_num(depth, nan=1, posinf=1, neginf=1), + 'depth_mask': np.isfinite(depth), + 'depth_mask_inf': np.isinf(depth), + }) + + if self.include_segmentation: + segmentation_mask, segmentation_labels = read_segmentation(Path(path,'segmentation.png')) + instance.update({ + 'segmentation_mask': segmentation_mask, + 'segmentation_labels': segmentation_labels, + }) + + meta = read_meta(Path(path, 'meta.json')) + instance['intrinsics'] = np.array(meta['intrinsics'], dtype=np.float32) + + return instance + + def _process_instance(self, instance: dict): + if instance is None: + return None + + image, depth, depth_mask, intrinsics = instance['image'], instance['depth'], instance['depth_mask'], instance['intrinsics'] + segmentation_mask, segmentation_labels = instance.get('segmentation_mask', None), instance.get('segmentation_labels', None) + + raw_height, raw_width = image.shape[:2] + raw_horizontal, raw_vertical = abs(1.0 / intrinsics[0, 0]), abs(1.0 / intrinsics[1, 1]) + raw_pixel_w, raw_pixel_h = raw_horizontal / raw_width, raw_vertical / raw_height + tgt_width, tgt_height = instance['width'], instance['height'] + tgt_aspect = tgt_width / tgt_height + + # set expected target view field + tgt_horizontal = min(raw_horizontal, raw_vertical * tgt_aspect) + tgt_vertical = tgt_horizontal / tgt_aspect + + # set target view direction + cu, cv = 0.5, 0.5 + direction = utils3d.np.unproject_cv(np.array([[cu, cv]], dtype=np.float32), np.array([1.0], dtype=np.float32), intrinsics=intrinsics)[0] + R = utils3d.np.rotation_matrix_from_vectors(direction, np.array([0, 0, 1], dtype=np.float32)) + + # restrict target view field within the raw view + corners = np.array([[0, 0], [0, 1], [1, 1], [1, 0]], dtype=np.float32) + corners = np.concatenate([corners, np.ones((4, 1), dtype=np.float32)], axis=1) @ (np.linalg.inv(intrinsics).T @ R.T) # corners in viewport's camera plane + corners = corners[:, :2] / corners[:, 2:3] + + warp_horizontal, warp_vertical = abs(1.0 / intrinsics[0, 0]), abs(1.0 / intrinsics[1, 1]) + for i in range(4): + intersection, _ = utils3d.np.ray_intersection( + np.array([0., 0.]), np.array([[tgt_aspect, 1.0], [tgt_aspect, -1.0]]), + corners[i - 1], corners[i] - corners[i - 1], + ) + warp_horizontal, warp_vertical = min(warp_horizontal, 2 * np.abs(intersection[:, 0]).min()), min(warp_vertical, 2 * np.abs(intersection[:, 1]).min()) + tgt_horizontal, tgt_vertical = min(tgt_horizontal, warp_horizontal), min(tgt_vertical, warp_vertical) + + # get target view intrinsics + fx, fy = 1.0 / tgt_horizontal, 1.0 / tgt_vertical + tgt_intrinsics = utils3d.np.intrinsics_from_focal_center(fx, fy, 0.5, 0.5).astype(np.float32) + + # do homogeneous transformation with the rotation and intrinsics + # 4.1 The image and depth is resized first to approximately the same pixel size as the target image with PIL's antialiasing resampling + tgt_pixel_w, tgt_pixel_h = tgt_horizontal / tgt_width, tgt_vertical / tgt_height # (should be exactly the same for x and y axes) + rescaled_w, rescaled_h = int(raw_width * raw_pixel_w / tgt_pixel_w), int(raw_height * raw_pixel_h / tgt_pixel_h) + image = np.array(Image.fromarray(image).resize((rescaled_w, rescaled_h), Image.Resampling.LANCZOS)) + + depth, depth_mask = utils3d.np.masked_nearest_resize(depth, mask=depth_mask, size=(rescaled_h, rescaled_w)) + distance = norm3d(utils3d.np.depth_map_to_point_map(depth, intrinsics=intrinsics)) + segmentation_mask = cv2.resize(segmentation_mask, (rescaled_w, rescaled_h), interpolation=cv2.INTER_NEAREST) if segmentation_mask is not None else None + + # 4.2 calculate homography warping + transform = intrinsics @ np.linalg.inv(R) @ np.linalg.inv(tgt_intrinsics) + uv_tgt = utils3d.np.uv_map(tgt_height, tgt_width) + pts = np.concatenate([uv_tgt, np.ones((tgt_height, tgt_width, 1), dtype=np.float32)], axis=-1) @ transform.T + uv_remap = pts[:, :, :2] / (pts[:, :, 2:3] + 1e-12) + pixel_remap = utils3d.np.uv_to_pixel(uv_remap, (rescaled_h, rescaled_w)).astype(np.float32) + + tgt_image = cv2.remap(image, pixel_remap[:, :, 0], pixel_remap[:, :, 1], cv2.INTER_LINEAR) + tgt_distance = cv2.remap(distance, pixel_remap[:, :, 0], pixel_remap[:, :, 1], cv2.INTER_NEAREST) + tgt_ray_length = utils3d.np.unproject_cv(uv_tgt, np.ones_like(uv_tgt[:, :, 0]), intrinsics=tgt_intrinsics) + tgt_ray_length = (tgt_ray_length[:, :, 0] ** 2 + tgt_ray_length[:, :, 1] ** 2 + tgt_ray_length[:, :, 2] ** 2) ** 0.5 + tgt_depth = tgt_distance / (tgt_ray_length + 1e-12) + tgt_depth_mask = cv2.remap(depth_mask.astype(np.uint8), pixel_remap[:, :, 0], pixel_remap[:, :, 1], cv2.INTER_NEAREST) > 0 + tgt_segmentation_mask = cv2.remap(segmentation_mask, pixel_remap[:, :, 0], pixel_remap[:, :, 1], cv2.INTER_NEAREST) if segmentation_mask is not None else None + + # drop depth greater than drop_max_depth + max_depth = np.nanquantile(np.where(tgt_depth_mask, tgt_depth, np.nan), 0.01) * self.drop_max_depth + tgt_depth_mask &= tgt_depth <= max_depth + tgt_depth = np.nan_to_num(tgt_depth, nan=0.0) + + if self.depth_unit is not None: + tgt_depth *= self.depth_unit + + if not np.any(tgt_depth_mask): + # always make sure that mask is not empty, otherwise the loss calculation will crash + tgt_depth_mask = np.ones_like(tgt_depth_mask) + tgt_depth = np.ones_like(tgt_depth) + instance['label_type'] = 'invalid' + + tgt_pts = utils3d.np.unproject_cv(uv_tgt, tgt_depth, intrinsics=tgt_intrinsics) + + # Process segmentation labels + if self.include_segmentation and segmentation_mask is not None: + for k in ['undefined', 'unannotated', 'background', 'sky']: + if k in segmentation_labels: + del segmentation_labels[k] + seg_id2count = dict(zip(*np.unique(tgt_segmentation_mask, return_counts=True))) + sorted_labels = sorted(segmentation_labels.keys(), key=lambda x: seg_id2count.get(segmentation_labels[x], 0), reverse=True) + segmentation_labels = {k: segmentation_labels[k] for k in sorted_labels[:self.max_segments] if seg_id2count.get(segmentation_labels[k], 0) >= self.min_seg_area} + + instance.update({ + 'image': torch.from_numpy(tgt_image.astype(np.float32) / 255.0).permute(2, 0, 1), + 'depth': torch.from_numpy(tgt_depth).float(), + 'depth_mask': torch.from_numpy(tgt_depth_mask).bool(), + 'intrinsics': torch.from_numpy(tgt_intrinsics).float(), + 'points': torch.from_numpy(tgt_pts).float(), + 'segmentation_mask': torch.from_numpy(tgt_segmentation_mask).long() if tgt_segmentation_mask is not None else None, + 'segmentation_labels': segmentation_labels, + 'is_metric': self.depth_unit is not None, + 'has_sharp_boundary': self.has_sharp_boundary, + }) + + instance = {k: v for k, v in instance.items() if v is not None} + + return instance + + def start(self): + self.pipeline.start() + + def stop(self): + self.pipeline.stop() + + def __enter__(self): + self.start() + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.stop() + + def get(self): + return self.pipeline.get() \ No newline at end of file diff --git a/moge/test/metrics.py b/moge/test/metrics.py new file mode 100644 index 0000000..4c79c33 --- /dev/null +++ b/moge/test/metrics.py @@ -0,0 +1,342 @@ +from typing import * +from numbers import Number + +import torch +import torch.nn.functional as F +import numpy as np +import utils3d + +from ..utils.geometry_torch import ( + weighted_mean, + intrinsics_to_fov +) +from ..utils.alignment import ( + align_points_scale_z_shift, + align_points_scale_xyz_shift, + align_points_xyz_shift, + align_affine_lstsq, + align_depth_scale, + align_depth_affine, + align_points_scale, +) +from ..utils.tools import key_average, timeit + + +def rel_depth(pred: torch.Tensor, gt: torch.Tensor, eps: float = 1e-6): + rel = (torch.abs(pred - gt) / (gt + eps)).mean() + return rel.item() + + +def delta1_depth(pred: torch.Tensor, gt: torch.Tensor, eps: float = 1e-6): + delta1 = (torch.maximum(gt / pred, pred / gt) < 1.25).float().mean() + return delta1.item() + + +def rel_point(pred: torch.Tensor, gt: torch.Tensor, eps: float = 1e-6): + dist_gt = torch.norm(gt, dim=-1) + dist_err = torch.norm(pred - gt, dim=-1) + rel = (dist_err / (dist_gt + eps)).mean() + return rel.item() + + +def delta1_point(pred: torch.Tensor, gt: torch.Tensor, eps: float = 1e-6): + dist_pred = torch.norm(pred, dim=-1) + dist_gt = torch.norm(gt, dim=-1) + dist_err = torch.norm(pred - gt, dim=-1) + + delta1 = (dist_err < 0.25 * torch.minimum(dist_gt, dist_pred)).float().mean() + return delta1.item() + + +def rel_point_local(pred: torch.Tensor, gt: torch.Tensor, diameter: torch.Tensor): + dist_err = torch.norm(pred - gt, dim=-1) + rel = (dist_err / diameter).mean() + return rel.item() + + +def delta1_point_local(pred: torch.Tensor, gt: torch.Tensor, diameter: torch.Tensor): + dist_err = torch.norm(pred - gt, dim=-1) + delta1 = (dist_err < 0.25 * diameter).float().mean() + return delta1.item() + + +def boundary_f1(pred: torch.Tensor, gt: torch.Tensor, mask: torch.Tensor, radius: int = 1): + neighbor_x, neight_y = torch.meshgrid( + torch.linspace(-radius, radius, 2 * radius + 1, device=pred.device), + torch.linspace(-radius, radius, 2 * radius + 1, device=pred.device), + indexing='xy' + ) + neighbor_mask = (neighbor_x ** 2 + neight_y ** 2) <= radius ** 2 + 1e-5 + + pred_window = utils3d.pt.sliding_window_2d(pred, window_size=2 * radius + 1, stride=1, dim=(-2, -1)) # [H, W, 2*R+1, 2*R+1] + gt_window = utils3d.pt.sliding_window_2d(gt, window_size=2 * radius + 1, stride=1, dim=(-2, -1)) # [H, W, 2*R+1, 2*R+1] + mask_window = neighbor_mask & utils3d.pt.sliding_window_2d(mask, window_size=2 * radius + 1, stride=1, dim=(-2, -1)) # [H, W, 2*R+1, 2*R+1] + + pred_rel = pred_window / pred[radius:-radius, radius:-radius, None, None] + gt_rel = gt_window / gt[radius:-radius, radius:-radius, None, None] + valid = mask[radius:-radius, radius:-radius, None, None] & mask_window + + f1_list = [] + w_list = t_list = torch.linspace(0.05, 0.25, 10).tolist() + + for t in t_list: + pred_label = pred_rel > 1 + t + gt_label = gt_rel > 1 + t + TP = (pred_label & gt_label & valid).float().sum() + precision = TP / (gt_label & valid).float().sum().clamp_min(1e-12) + recall = TP / (pred_label & valid).float().sum().clamp_min(1e-12) + f1 = 2 * precision * recall / (precision + recall).clamp_min(1e-12) + f1_list.append(f1.item()) + + f1_avg = sum(w * f1 for w, f1 in zip(w_list, f1_list)) / sum(w_list) + return f1_avg + + +def compute_metrics( + pred: Dict[str, torch.Tensor], + gt: Dict[str, torch.Tensor], + vis: bool = False +) -> Tuple[Dict[str, Dict[str, Number]], Dict[str, torch.Tensor]]: + """ + A unified function to compute metrics for different types of predictions and ground truths. + + #### Supported keys in pred: + - `disparity_affine_invariant`: disparity map predicted by a depth estimator with scale and shift invariant. + - `depth_scale_invariant`: depth map predicted by a depth estimator with scale invariant. + - `depth_affine_invariant`: depth map predicted by a depth estimator with scale and shift invariant. + - `depth_metric`: depth map predicted by a depth estimator with no scale or shift. + - `points_scale_invariant`: point map predicted by a point estimator with scale invariant. + - `points_affine_invariant`: point map predicted by a point estimator with scale and xyz shift invariant. + - `points_metric`: point map predicted by a point estimator with no scale or shift. + - `intrinsics`: normalized camera intrinsics matrix. + + #### Required keys in gt: + - `depth`: depth map ground truth (in metric units if `depth_metric` is used) + - `points`: point map ground truth in camera coordinates. + - `mask`: mask indicating valid pixels in the ground truth. + - `intrinsics`: normalized ground-truth camera intrinsics matrix. + - `is_metric`: whether the depth is in metric units. + """ + metrics = {} + misc = {} + + mask = gt['depth_mask'] + gt_depth = gt['depth'] + gt_points = gt['points'] + + height, width = mask.shape[-2:] + lr_mask, lr_index = utils3d.pt.masked_nearest_resize(mask=mask, size=(64, 64), return_index=True) + + only_depth = not any('point' in k for k in pred) + pred_depth_aligned, pred_points_aligned = None, None + + # Metric depth + if 'depth_metric' in pred and gt['is_metric']: + pred_depth, gt_depth = pred['depth_metric'], gt['depth'] + metrics['depth_metric'] = { + 'rel': rel_depth(pred_depth[mask], gt_depth[mask]), + 'delta1': delta1_depth(pred_depth[mask], gt_depth[mask]) + } + + if pred_depth_aligned is None: + pred_depth_aligned = pred_depth + + # Scale-invariant depth + if 'depth_scale_invariant' in pred: + pred_depth_scale_invariant = pred['depth_scale_invariant'] + elif 'depth_metric' in pred: + pred_depth_scale_invariant = pred['depth_metric'] + else: + pred_depth_scale_invariant = None + + if pred_depth_scale_invariant is not None: + pred_depth = pred_depth_scale_invariant + + pred_depth_lr_masked, gt_depth_lr_masked = pred_depth[lr_index][lr_mask], gt_depth[lr_index][lr_mask] + scale = align_depth_scale(pred_depth_lr_masked, gt_depth_lr_masked, 1 / gt_depth_lr_masked) + pred_depth = pred_depth * scale + + metrics['depth_scale_invariant'] = { + 'rel': rel_depth(pred_depth[mask], gt_depth[mask]), + 'delta1': delta1_depth(pred_depth[mask], gt_depth[mask]) + } + + if pred_depth_aligned is None: + pred_depth_aligned = pred_depth + + # Affine-invariant depth + if 'depth_affine_invariant' in pred: + pred_depth_affine_invariant = pred['depth_affine_invariant'] + elif 'depth_scale_invariant' in pred: + pred_depth_affine_invariant = pred['depth_scale_invariant'] + elif 'depth_metric' in pred: + pred_depth_affine_invariant = pred['depth_metric'] + else: + pred_depth_affine_invariant = None + + if pred_depth_affine_invariant is not None: + pred_depth = pred_depth_affine_invariant + + pred_depth_lr_masked, gt_depth_lr_masked = pred_depth[lr_index][lr_mask], gt_depth[lr_index][lr_mask] + scale, shift = align_depth_affine(pred_depth_lr_masked, gt_depth_lr_masked, 1 / gt_depth_lr_masked) + pred_depth = pred_depth * scale + shift + + metrics['depth_affine_invariant'] = { + 'rel': rel_depth(pred_depth[mask], gt_depth[mask]), + 'delta1': delta1_depth(pred_depth[mask], gt_depth[mask]) + } + + if pred_depth_aligned is None: + pred_depth_aligned = pred_depth + + # Affine-invariant disparity + if 'disparity_affine_invariant' in pred: + pred_disparity_affine_invariant = pred['disparity_affine_invariant'] + elif 'depth_scale_invariant' in pred: + pred_disparity_affine_invariant = 1 / pred['depth_scale_invariant'] + elif 'depth_metric' in pred: + pred_disparity_affine_invariant = 1 / pred['depth_metric'] + else: + pred_disparity_affine_invariant = None + + if pred_disparity_affine_invariant is not None: + pred_disp = pred_disparity_affine_invariant + + scale, shift = align_affine_lstsq(pred_disp[mask], 1 / gt_depth[mask]) + pred_disp = pred_disp * scale + shift + + # NOTE: The alignment is done on the disparity map could introduce extreme outliers at disparities close to 0. + # Therefore we clamp the disparities by minimum ground truth disparity. + pred_depth = 1 / pred_disp.clamp_min(1 / gt_depth[mask].max().item()) + + metrics['disparity_affine_invariant'] = { + 'rel': rel_depth(pred_depth[mask], gt_depth[mask]), + 'delta1': delta1_depth(pred_depth[mask], gt_depth[mask]) + } + + if pred_depth_aligned is None: + pred_depth_aligned = 1 / pred_disp.clamp_min(1e-6) + + # Metric points + if 'points_metric' in pred and gt['is_metric']: + pred_points = pred['points_metric'] + + pred_points_lr_masked, gt_points_lr_masked = pred_points[lr_index][lr_mask], gt_points[lr_index][lr_mask] + shift = align_points_xyz_shift(pred_points_lr_masked, gt_points_lr_masked, 1 / gt_points_lr_masked.norm(dim=-1)) + pred_points = pred_points + shift + + metrics['points_metric'] = { + 'rel': rel_point(pred_points[mask], gt_points[mask]), + 'delta1': delta1_point(pred_points[mask], gt_points[mask]) + } + + if pred_points_aligned is None: + pred_points_aligned = pred['points_metric'] + + # Scale-invariant points (in camera space) + if 'points_scale_invariant' in pred: + pred_points_scale_invariant = pred['points_scale_invariant'] + elif 'points_metric' in pred: + pred_points_scale_invariant = pred['points_metric'] + else: + pred_points_scale_invariant = None + + if pred_points_scale_invariant is not None: + pred_points = pred_points_scale_invariant + + pred_points_lr_masked, gt_points_lr_masked = pred_points_scale_invariant[lr_index][lr_mask], gt_points[lr_index][lr_mask] + scale = align_points_scale(pred_points_lr_masked, gt_points_lr_masked, 1 / gt_points_lr_masked.norm(dim=-1)) + pred_points = pred_points * scale + + metrics['points_scale_invariant'] = { + 'rel': rel_point(pred_points[mask], gt_points[mask]), + 'delta1': delta1_point(pred_points[mask], gt_points[mask]) + } + + if vis and pred_points_aligned is None: + pred_points_aligned = pred['points_scale_invariant'] * scale + + # Affine-invariant points + if 'points_affine_invariant' in pred: + pred_points_affine_invariant = pred['points_affine_invariant'] + elif 'points_scale_invariant' in pred: + pred_points_affine_invariant = pred['points_scale_invariant'] + elif 'points_metric' in pred: + pred_points_affine_invariant = pred['points_metric'] + else: + pred_points_affine_invariant = None + + if pred_points_affine_invariant is not None: + pred_points = pred_points_affine_invariant + + pred_points_lr_masked, gt_points_lr_masked = pred_points[lr_index][lr_mask], gt_points[lr_index][lr_mask] + scale, shift = align_points_scale_xyz_shift(pred_points_lr_masked, gt_points_lr_masked, 1 / gt_points_lr_masked.norm(dim=-1)) + pred_points = pred_points * scale + shift + + metrics['points_affine_invariant'] = { + 'rel': rel_point(pred_points[mask], gt_points[mask]), + 'delta1': delta1_point(pred_points[mask], gt_points[mask]) + } + + if vis and pred_points_aligned is None: + pred_points_aligned = pred['points_affine_invariant'] * scale + shift + + # Local points + if 'segmentation_mask' in gt and 'points' in gt and any('points' in k for k in pred.keys()): + pred_points = next(pred[k] for k in pred.keys() if 'points' in k) + gt_points = gt['points'] + segmentation_mask = gt['segmentation_mask'] + segmentation_labels = gt['segmentation_labels'] + segmentation_mask_lr = segmentation_mask[lr_index] + local_points_metrics = [] + for _, seg_id in segmentation_labels.items(): + valid_mask = (segmentation_mask == seg_id) & mask + + pred_points_masked = pred_points[valid_mask] + gt_points_masked = gt_points[valid_mask] + + valid_mask_lr = (segmentation_mask_lr == seg_id) & lr_mask + if valid_mask_lr.sum().item() < 10: + continue + pred_points_masked_lr = pred_points[lr_index][valid_mask_lr] + gt_points_masked_lr = gt_points[lr_index][valid_mask_lr] + diameter = (gt_points_masked.max(dim=0).values - gt_points_masked.min(dim=0).values).max() + scale, shift = align_points_scale_xyz_shift(pred_points_masked_lr, gt_points_masked_lr, 1 / diameter.expand(gt_points_masked_lr.shape[0])) + pred_points_masked = pred_points_masked * scale + shift + + local_points_metrics.append({ + 'rel': rel_point_local(pred_points_masked, gt_points_masked, diameter), + 'delta1': delta1_point_local(pred_points_masked, gt_points_masked, diameter), + }) + + metrics['local_points'] = key_average(local_points_metrics) + + # FOV. NOTE: If there is no random augmentation applied to the input images, all GT FOV are generallly the same. + # Fair evaluation of FOV requires random augmentation. + if 'intrinsics' in pred and 'intrinsics' in gt: + pred_intrinsics = pred['intrinsics'] + gt_intrinsics = gt['intrinsics'] + pred_fov_x, pred_fov_y = intrinsics_to_fov(pred_intrinsics) + gt_fov_x, gt_fov_y = intrinsics_to_fov(gt_intrinsics) + metrics['fov_x'] = { + 'mae': torch.rad2deg(pred_fov_x - gt_fov_x).abs().mean().item(), + 'deviation': torch.rad2deg(pred_fov_x - gt_fov_x).item(), + } + + # Boundary F1 + if pred_depth_aligned is not None and gt['has_sharp_boundary']: + metrics['boundary'] = { + 'radius1_f1': boundary_f1(pred_depth_aligned, gt_depth, mask, radius=1), + 'radius2_f1': boundary_f1(pred_depth_aligned, gt_depth, mask, radius=2), + 'radius3_f1': boundary_f1(pred_depth_aligned, gt_depth, mask, radius=3), + } + + if vis: + if pred_points_aligned is not None: + misc['pred_points'] = pred_points_aligned + if only_depth: + misc['pred_points'] = utils3d.pt.depth_map_to_point_map(pred_depth_aligned, intrinsics=gt['intrinsics']) + if pred_depth_aligned is not None: + misc['pred_depth'] = pred_depth_aligned + + return metrics, misc \ No newline at end of file diff --git a/moge/train/__init__.py b/moge/train/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/moge/train/dataloader.py b/moge/train/dataloader.py new file mode 100644 index 0000000..b08846d --- /dev/null +++ b/moge/train/dataloader.py @@ -0,0 +1,258 @@ +import os +from pathlib import Path +import json +import time +import random +from typing import * +import traceback +import itertools +from numbers import Number +import io + +import numpy as np +import cv2 +from PIL import Image +import torch +import torchvision.transforms.v2.functional as TF +import utils3d +import pipeline +from tqdm import tqdm + +from ..utils.io import * +from ..utils.geometry_numpy import harmonic_mean_numpy, norm3d, depth_occlusion_edge_numpy +from ..utils.data_augmentation import sample_perspective, warp_perspective, image_color_augmentation + + +class TrainDataLoaderPipeline: + def __init__(self, config: dict, batch_size: int, num_load_workers: int = 4, num_process_workers: int = 8, buffer_size: int = 8): + self.config = config + + self.batch_size = batch_size + self.clamp_max_depth = config['clamp_max_depth'] + self.fov_range_absolute = config.get('fov_range_absolute', 0.0) + self.fov_range_relative = config.get('fov_range_relative', 0.0) + self.center_augmentation = config.get('center_augmentation', 0.0) + self.image_augmentation = config.get('image_augmentation', []) + self.depth_interpolation = config.get('depth_interpolation', 'bilinear') + + if 'image_sizes' in config: + self.image_size_strategy = 'fixed' + self.image_sizes = config['image_sizes'] + elif 'aspect_ratio_range' in config and 'area_range' in config: + self.image_size_strategy = 'aspect_area' + self.aspect_ratio_range = config['aspect_ratio_range'] + self.area_range = config['area_range'] + else: + raise ValueError('Invalid image size configuration') + + # Load datasets + self.datasets = {} + for dataset in tqdm(config['datasets'], desc='Loading datasets'): + name = dataset['name'] + content = Path(dataset['path'], dataset.get('index', '.index.txt')).joinpath().read_text() + filenames = content.splitlines() + self.datasets[name] = { + **dataset, + 'path': dataset['path'], + 'filenames': filenames, + } + self.dataset_names = [dataset['name'] for dataset in config['datasets']] + self.dataset_weights = [dataset['weight'] for dataset in config['datasets']] + + # Build pipeline + self.pipeline = pipeline.Sequential([ + self._sample_batch, + pipeline.Unbatch(), + pipeline.Parallel([self._load_instance] * num_load_workers), + pipeline.Parallel([self._process_instance] * num_process_workers), + pipeline.Batch(self.batch_size), + self._collate_batch, + pipeline.Buffer(buffer_size), + ]) + + self.invalid_instance = { + 'intrinsics': np.array([[1.0, 0.0, 0.5], [0.0, 1.0, 0.5], [0.0, 0.0, 1.0]], dtype=np.float32), + 'image': np.zeros((256, 256, 3), dtype=np.uint8), + 'depth': np.ones((256, 256), dtype=np.float32), + 'depth_mask': np.ones((256, 256), dtype=bool), + 'depth_mask_inf': np.zeros((256, 256), dtype=bool), + 'label_type': 'invalid', + } + + def _sample_batch(self): + batch_id = 0 + last_area = None + while True: + # Depending on the sample strategy, choose a dataset and a filename + batch_id += 1 + batch = [] + + # Sample instances + for _ in range(self.batch_size): + dataset_name = random.choices(self.dataset_names, weights=self.dataset_weights)[0] + filename = random.choice(self.datasets[dataset_name]['filenames']) + + path = Path(self.datasets[dataset_name]['path'], filename) + + instance = { + 'batch_id': batch_id, + 'seed': random.randint(0, 2 ** 32 - 1), + 'dataset': dataset_name, + 'filename': filename, + 'path': path, + 'label_type': self.datasets[dataset_name]['label_type'], + } + batch.append(instance) + + # Decide the image size for this batch + if self.image_size_strategy == 'fixed': + width, height = random.choice(self.config['image_sizes']) + elif self.image_size_strategy == 'aspect_area': + area = random.uniform(*self.area_range) + aspect_ratio_ranges = [self.datasets[instance['dataset']].get('aspect_ratio_range', self.aspect_ratio_range) for instance in batch] + aspect_ratio_range = (min(r[0] for r in aspect_ratio_ranges), max(r[1] for r in aspect_ratio_ranges)) + aspect_ratio = random.uniform(*aspect_ratio_range) + width, height = int((area * aspect_ratio) ** 0.5), int((area / aspect_ratio) ** 0.5) + else: + raise ValueError('Invalid image size strategy') + + for instance in batch: + instance['width'], instance['height'] = width, height + + yield batch + + def _load_instance(self, instance: dict): + try: + image = read_image(Path(instance['path'], 'image.jpg')) + depth = read_depth(Path(instance['path'], self.datasets[instance['dataset']].get('depth', 'depth.png'))) + meta = read_json(Path(instance['path'], 'meta.json')) + intrinsics = np.array(meta['intrinsics'], dtype=np.float32) + data = { + 'image': image, + 'depth': depth, + 'intrinsics': intrinsics + } + instance.update({ + **data, + }) + except Exception as e: + print(f"Failed to load instance {instance['dataset']}/{instance['filename']} because of exception:", e) + instance.update(self.invalid_instance) + return instance + + def _process_instance(self, instance: Dict[str, Union[np.ndarray, str, float, bool]]): + raw_image, raw_depth, raw_intrinsics, label_type = instance['image'], instance['depth'], instance['intrinsics'], instance['label_type'] + raw_normal, raw_normal_mask = utils3d.np.depth_map_to_normal_map(raw_depth, intrinsics=raw_intrinsics, mask=np.isfinite(raw_depth), edge_threshold=88) + raw_normal = np.where(raw_normal_mask[..., None], raw_normal, np.nan) + depth_unit = self.datasets[instance['dataset']].get('depth_unit', None) + + raw_height, raw_width = raw_image.shape[:2] + raw_fov_x, raw_fov_y = utils3d.np.intrinsics_to_fov(raw_intrinsics) + tgt_width, tgt_height = instance['width'], instance['height'] + tgt_aspect = tgt_width / tgt_height + + rng = np.random.default_rng(instance['seed']) + + # Sample perspective transformation + tgt_intrinsics, R = sample_perspective( + raw_intrinsics, + tgt_aspect=tgt_aspect, + center_augmentation=self.datasets[instance['dataset']].get('center_augmentation', self.center_augmentation), + fov_range_absolute=self.datasets[instance['dataset']].get('fov_range_absolute', self.fov_range_absolute), + fov_range_relative=self.datasets[instance['dataset']].get('fov_range_relative', self.fov_range_relative), + rng=rng + ) + + # Warp + transform = tgt_intrinsics @ R @ np.linalg.inv(raw_intrinsics) + # - Warp image + tgt_image = warp_perspective(raw_image, transform, tgt_size=(tgt_height, tgt_width), interpolation='lanczos') + # - Warp depth + depth_edge_mask = utils3d.np.depth_map_edge(raw_depth, mask=np.isfinite(raw_depth), kernel_size=5, ltol=0.01) + depth_bilinear_mask = np.isfinite(raw_depth) & ~depth_edge_mask + warped_depth_bilinear_mask = warp_perspective(depth_bilinear_mask.astype(np.float32), transform, (tgt_height, tgt_width), interpolation='bilinear') + warped_depth_nearest = warp_perspective(raw_depth, transform, (tgt_height, tgt_width), interpolation='nearest', sparse_mask=~np.isnan(raw_depth)) + warped_depth_bilinear = 1 / warp_perspective(1 / raw_depth, transform, (tgt_height, tgt_width), interpolation='bilinear') # NOTE: Bilinear intepolation in disparity space maintains planar surfaces. + warped_depth = np.where(warped_depth_bilinear_mask == 1., warped_depth_bilinear, warped_depth_nearest) + tgt_uvhomo = np.concatenate([utils3d.np.uv_map((tgt_height, tgt_width)), np.ones((tgt_height, tgt_width, 1), dtype=np.float32)], axis=-1) + tgt_depth = warped_depth / np.dot(tgt_uvhomo, np.linalg.inv(transform)[2, :]) + # - Warp normal + warped_normal = warp_perspective(raw_normal, transform, (tgt_height, tgt_width), interpolation='bilinear') + tgt_normal = warped_normal @ R.T + + # always make sure that mask is not empty + if np.isfinite(tgt_depth).sum() / tgt_depth.size < 0.001: + tgt_depth = np.ones_like(tgt_depth) + instance['label_type'] = 'invalid' + + # Flip augmentation + if rng.choice([True, False]): + tgt_image = np.flip(tgt_image, axis=1).copy() + tgt_depth = np.flip(tgt_depth, axis=1).copy() + tgt_normal = np.flip(tgt_normal, axis=1).copy() * [-1, 1, 1] + # NOTE: if cx != 0.5, flip intrinsics accordingly. + + # Color augmentation + image_augmentation = self.datasets[instance['dataset']].get('image_augmentation', self.image_augmentation) + tgt_image = image_color_augmentation( + tgt_image, + augmentations=image_augmentation, + rng=rng, + depth=tgt_depth, + ) + + # Set metric flag if depth is in metric unit + if depth_unit is not None: + tgt_depth *= depth_unit + instance['is_metric'] = True + else: + instance['is_metric'] = False + + # Clip maximum depth + max_depth = np.nanquantile(np.where(np.isfinite(tgt_depth), tgt_depth, np.nan), 0.01) * self.clamp_max_depth + tgt_depth = np.where(np.isfinite(tgt_depth), np.clip(tgt_depth, 0, max_depth), tgt_depth) + + tgt_depth_mask_inf = np.isinf(tgt_depth) + if self.datasets[instance['dataset']].get('finite_depth_mask', None) == "only_known": + tgt_depth_mask_fin = np.isfinite(tgt_depth) + else: + tgt_depth_mask_fin = ~tgt_depth_mask_inf + + instance.update({ + 'image': torch.from_numpy(tgt_image.astype(np.float32) / 255.0).permute(2, 0, 1), + 'depth': torch.from_numpy(tgt_depth).float(), + 'depth_mask_fin': torch.from_numpy(tgt_depth_mask_fin).bool(), + 'depth_mask_inf': torch.from_numpy(tgt_depth_mask_inf).bool(), + "normal": torch.from_numpy(tgt_normal).float(), + 'intrinsics': torch.from_numpy(tgt_intrinsics).float(), + }) + return instance + + def _collate_batch(self, instances: List[Dict[str, Any]]): + batch = {k: torch.stack([instance[k] for instance in instances], dim=0) for k in ['image', 'depth', 'depth_mask_fin', 'depth_mask_inf', 'normal', 'intrinsics']} + batch = { + 'label_type': [instance['label_type'] for instance in instances], + 'is_metric': [instance['is_metric'] for instance in instances], + 'info': [{'dataset': instance['dataset'], 'filename': instance['filename']} for instance in instances], + **batch, + } + return batch + + def get(self) -> Dict[str, Union[torch.Tensor, str]]: + return self.pipeline.get() + + def start(self): + self.pipeline.start() + + def stop(self): + self.pipeline.stop() + + def __enter__(self): + self.start() + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.pipeline.stop() + return False + + diff --git a/moge/train/losses.py b/moge/train/losses.py new file mode 100644 index 0000000..a568adf --- /dev/null +++ b/moge/train/losses.py @@ -0,0 +1,293 @@ +from typing import * +import math + +import torch +import torch.nn.functional as F +import utils3d + +from ..utils.geometry_torch import ( + weighted_mean, + harmonic_mean, + geometric_mean, + normalized_view_plane_uv, + angle_diff_vec3 +) +from ..utils.alignment import ( + align_points_scale_z_shift, + align_points_scale, + align_points_scale_xyz_shift, + align_points_z_shift, +) + + +def _smooth(err: torch.FloatTensor, beta: float = 0.0) -> torch.FloatTensor: + if beta == 0: + return err + else: + return torch.where(err < beta, 0.5 * err.square() / beta, err - 0.5 * beta) + + +def affine_invariant_global_loss( + pred_points: torch.Tensor, + gt_points: torch.Tensor, + align_resolution: int = 64, + beta: float = 0.0, + trunc: float = 1.0, + sparsity_aware: bool = False +): + device = pred_points.device + + mask = torch.isfinite(gt_points).all(dim=-1) + gt_points = torch.where(mask[..., None], gt_points, 1) + + # Align + pred_points_lr, gt_points_lr, lr_mask = utils3d.pt.masked_nearest_resize(pred_points, gt_points, mask=mask, size=(align_resolution, align_resolution)) + scale, shift = align_points_scale_z_shift(pred_points_lr.flatten(-3, -2), gt_points_lr.flatten(-3, -2), lr_mask.flatten(-2, -1) / gt_points_lr[..., 2].flatten(-2, -1).clamp_min(1e-2), trunc=trunc) + valid = scale > 0 + scale, shift = torch.where(valid, scale, 0), torch.where(valid[..., None], shift, 0) + + pred_points = scale[..., None, None, None] * pred_points + shift[..., None, None, :] + + # Compute loss + weight = (valid[..., None, None] & mask).float() / gt_points[..., 2].clamp_min(1e-5) + weight = weight.clamp_max(10.0 * weighted_mean(weight, mask, dim=(-2, -1), keepdim=True)) # In case your data contains extremely small depth values + loss = _smooth((pred_points - gt_points).abs() * weight[..., None], beta=beta).mean(dim=(-3, -2, -1)) + + if sparsity_aware: + # Reweighting improves performance on sparse depth data. NOTE: this is not used in MoGe-1. + sparsity = mask.float().mean(dim=(-2, -1)) / lr_mask.float().mean(dim=(-2, -1)) + loss = loss / (sparsity + 1e-7) + + err = (pred_points.detach() - gt_points).norm(dim=-1) / gt_points[..., 2] + + # Record any scalar metric + misc = { + 'truncated_error': weighted_mean(err.clamp_max(1.0), mask).item(), + 'delta': weighted_mean((err < 1).float(), mask).item() + } + + return loss, misc, scale.detach() + + +def monitoring(points: torch.Tensor): + return { + 'std': points.std().item(), + } + + +def compute_anchor_sampling_weight( + points: torch.Tensor, + mask: torch.Tensor, + radius_2d: torch.Tensor, + radius_3d: torch.Tensor, + num_test: int = 64 +) -> torch.Tensor: + # Importance sampling to balance the sampled probability of fine strutures. + # NOTE: MoGe-1 uses uniform random sampling instead of importance sampling. + # This is an incremental trick introduced later than the publication of MoGe-1 paper. + + height, width = points.shape[-3:-1] + + pixel_i, pixel_j = torch.meshgrid( + torch.arange(height, device=points.device), + torch.arange(width, device=points.device), + indexing='ij' + ) + + test_delta_i = torch.randint(-radius_2d, radius_2d + 1, (height, width, num_test,), device=points.device) # [num_test] + test_delta_j = torch.randint(-radius_2d, radius_2d + 1, (height, width, num_test,), device=points.device) # [num_test] + test_i, test_j = pixel_i[..., None] + test_delta_i, pixel_j[..., None] + test_delta_j # [height, width, num_test] + test_mask = (test_i >= 0) & (test_i < height) & (test_j >= 0) & (test_j < width) # [height, width, num_test] + test_i, test_j = test_i.clamp(0, height - 1), test_j.clamp(0, width - 1) # [height, width, num_test] + test_mask = test_mask & mask[..., test_i, test_j] # [..., height, width, num_test] + test_points = points[..., test_i, test_j, :] # [..., height, width, num_test, 3] + test_dist = (test_points - points[..., None, :]).norm(dim=-1) # [..., height, width, num_test] + + weight = 1 / ((test_dist <= radius_3d[..., None]) & test_mask).float().sum(dim=-1).clamp_min(1) + weight = torch.where(mask, weight, 0) + weight = weight / weight.sum(dim=(-2, -1), keepdim=True).add(1e-7) # [..., height, width] + return weight + + +def affine_invariant_local_loss( + pred_points: torch.Tensor, + gt_points: torch.Tensor, + focal: torch.Tensor, + global_scale: torch.Tensor, + level: Literal[4, 16, 64], + align_resolution: int = 32, + num_patches: int = 16, + beta: float = 0.0, + trunc: float = 1.0, + sparsity_aware: bool = False +): + device, dtype = pred_points.device, pred_points.dtype + *batch_shape, height, width, _ = pred_points.shape + batch_size = math.prod(batch_shape) + + gt_mask = torch.isfinite(gt_points).all(dim=-1) + gt_points = torch.where(gt_mask[..., None], gt_points, 1) + pred_points, gt_points, gt_mask, focal, global_scale = pred_points.reshape(-1, height, width, 3), gt_points.reshape(-1, height, width, 3), gt_mask.reshape(-1, height, width), focal.reshape(-1), global_scale.reshape(-1) if global_scale is not None else None + + # Sample patch anchor points indices [num_total_patches] + radius_2d = math.ceil(0.5 / level * (height ** 2 + width ** 2) ** 0.5) + radius_3d = 0.5 / level / focal * gt_points[..., 2] + anchor_sampling_weights = compute_anchor_sampling_weight(gt_points, gt_mask, radius_2d, radius_3d, num_test=64) + where_mask = torch.where(gt_mask) + random_selection = torch.multinomial(anchor_sampling_weights[where_mask], num_patches * batch_size, replacement=True) + patch_batch_idx, patch_anchor_i, patch_anchor_j = [indices[random_selection] for indices in where_mask] # [num_total_patches] + + # Get patch indices [num_total_patches, patch_h, patch_w] + patch_i, patch_j = torch.meshgrid( + torch.arange(-radius_2d, radius_2d + 1, device=device), + torch.arange(-radius_2d, radius_2d + 1, device=device), + indexing='ij' + ) + patch_i, patch_j = patch_i + patch_anchor_i[:, None, None], patch_j + patch_anchor_j[:, None, None] + patch_mask = (patch_i >= 0) & (patch_i < height) & (patch_j >= 0) & (patch_j < width) + patch_i, patch_j = patch_i.clamp(0, height - 1), patch_j.clamp(0, width - 1) + + # Get patch mask and gt patch points + gt_patch_anchor_points = gt_points[patch_batch_idx, patch_anchor_i, patch_anchor_j] + gt_patch_radius_3d = 0.5 / level / focal[patch_batch_idx] * gt_patch_anchor_points[:, 2] + gt_patch_points = gt_points[patch_batch_idx[:, None, None], patch_i, patch_j] + gt_patch_dist = (gt_patch_points - gt_patch_anchor_points[:, None, None, :]).norm(dim=-1) + patch_mask &= gt_mask[patch_batch_idx[:, None, None], patch_i, patch_j] + patch_mask &= gt_patch_dist <= gt_patch_radius_3d[:, None, None] + + # Pick only non-empty patches + MINIMUM_POINTS_PER_PATCH = 32 + nonempty = torch.where(patch_mask.sum(dim=(-2, -1)) >= MINIMUM_POINTS_PER_PATCH) + num_nonempty_patches = nonempty[0].shape[0] + if num_nonempty_patches == 0: + return torch.tensor(0.0, dtype=dtype, device=device), {} + + # Finalize all patch variables + patch_batch_idx, patch_i, patch_j = patch_batch_idx[nonempty], patch_i[nonempty], patch_j[nonempty] + patch_mask = patch_mask[nonempty] # [num_nonempty_patches, patch_h, patch_w] + gt_patch_points = gt_patch_points[nonempty] # [num_nonempty_patches, patch_h, patch_w, 3] + gt_patch_radius_3d = gt_patch_radius_3d[nonempty] # [num_nonempty_patches] + gt_patch_anchor_points = gt_patch_anchor_points[nonempty] # [num_nonempty_patches, 3] + pred_patch_points = pred_points[patch_batch_idx[:, None, None], patch_i, patch_j] + + # Align patch points + pred_patch_points_lr, gt_patch_points_lr, patch_lr_mask = utils3d.pt.masked_nearest_resize(pred_patch_points, gt_patch_points, mask=patch_mask, size=(align_resolution, align_resolution)) + local_scale, local_shift = align_points_scale_xyz_shift(pred_patch_points_lr.flatten(-3, -2), gt_patch_points_lr.flatten(-3, -2), patch_lr_mask.flatten(-2) / gt_patch_radius_3d[:, None].add(1e-7), trunc=trunc) + if global_scale is not None: + scale_differ = local_scale / global_scale[patch_batch_idx] + patch_valid = (scale_differ > 0.1) & (scale_differ < 10.0) & (global_scale > 0) + else: + patch_valid = local_scale > 0 + local_scale, local_shift = torch.where(patch_valid, local_scale, 0), torch.where(patch_valid[:, None], local_shift, 0) + patch_mask &= patch_valid[:, None, None] + + pred_patch_points = local_scale[:, None, None, None] * pred_patch_points + local_shift[:, None, None, :] # [num_patches_nonempty, patch_h, patch_w, 3] + + # Compute loss + gt_mean = harmonic_mean(gt_points[..., 2], gt_mask, dim=(-2, -1)) + patch_weight = patch_mask.float() / gt_patch_points[..., 2].clamp_min(0.1 * gt_mean[patch_batch_idx, None, None]) # [num_patches_nonempty, patch_h, patch_w] + loss = _smooth((pred_patch_points - gt_patch_points).abs() * patch_weight[..., None], beta=beta).mean(dim=(-3, -2, -1)) # [num_patches_nonempty] + + if sparsity_aware: + # Reweighting improves performance on sparse depth data. NOTE: this is not used in MoGe-1. + sparsity = patch_mask.float().mean(dim=(-2, -1)) / patch_lr_mask.float().mean(dim=(-2, -1)) + loss = loss / (sparsity + 1e-7) + loss = torch.scatter_reduce(torch.zeros(batch_size, dtype=dtype, device=device), dim=0, index=patch_batch_idx, src=loss, reduce='sum') / num_patches + loss = loss.reshape(batch_shape) + + err = (pred_patch_points.detach() - gt_patch_points).norm(dim=-1) / gt_patch_radius_3d[..., None, None] + + # Record any scalar metric + misc = { + 'truncated_error': weighted_mean(err.clamp_max(1), patch_mask).item(), + 'delta': weighted_mean((err < 1).float(), patch_mask).item() + } + + return loss, misc + + +def normal_loss(points: torch.Tensor, gt_points: torch.Tensor) -> torch.Tensor: + device, dtype = points.device, points.dtype + height, width = points.shape[-3:-1] + + mask = torch.isfinite(gt_points).all(dim=-1) + gt_points = torch.where(mask[..., None], gt_points, 1) + + leftup, rightup, leftdown, rightdown = points[..., :-1, :-1, :], points[..., :-1, 1:, :], points[..., 1:, :-1, :], points[..., 1:, 1:, :] + upxleft = torch.cross(rightup - rightdown, leftdown - rightdown, dim=-1) + leftxdown = torch.cross(leftup - rightup, rightdown - rightup, dim=-1) + downxright = torch.cross(leftdown - leftup, rightup - leftup, dim=-1) + rightxup = torch.cross(rightdown - leftdown, leftup - leftdown, dim=-1) + + gt_leftup, gt_rightup, gt_leftdown, gt_rightdown = gt_points[..., :-1, :-1, :], gt_points[..., :-1, 1:, :], gt_points[..., 1:, :-1, :], gt_points[..., 1:, 1:, :] + gt_upxleft = torch.cross(gt_rightup - gt_rightdown, gt_leftdown - gt_rightdown, dim=-1) + gt_leftxdown = torch.cross(gt_leftup - gt_rightup, gt_rightdown - gt_rightup, dim=-1) + gt_downxright = torch.cross(gt_leftdown - gt_leftup, gt_rightup - gt_leftup, dim=-1) + gt_rightxup = torch.cross(gt_rightdown - gt_leftdown, gt_leftup - gt_leftdown, dim=-1) + + mask_leftup, mask_rightup, mask_leftdown, mask_rightdown = mask[..., :-1, :-1], mask[..., :-1, 1:], mask[..., 1:, :-1], mask[..., 1:, 1:] + mask_upxleft = mask_rightup & mask_leftdown & mask_rightdown + mask_leftxdown = mask_leftup & mask_rightdown & mask_rightup + mask_downxright = mask_leftdown & mask_rightup & mask_leftup + mask_rightxup = mask_rightdown & mask_leftup & mask_leftdown + + MIN_ANGLE, MAX_ANGLE, BETA_RAD = math.radians(1), math.radians(90), math.radians(3) + + loss = mask_upxleft * _smooth(angle_diff_vec3(upxleft, gt_upxleft).clamp(MIN_ANGLE, MAX_ANGLE), beta=BETA_RAD) \ + + mask_leftxdown * _smooth(angle_diff_vec3(leftxdown, gt_leftxdown).clamp(MIN_ANGLE, MAX_ANGLE), beta=BETA_RAD) \ + + mask_downxright * _smooth(angle_diff_vec3(downxright, gt_downxright).clamp(MIN_ANGLE, MAX_ANGLE), beta=BETA_RAD) \ + + mask_rightxup * _smooth(angle_diff_vec3(rightxup, gt_rightxup).clamp(MIN_ANGLE, MAX_ANGLE), beta=BETA_RAD) + + loss = loss.mean() / (4 * max(points.shape[-3:-1])) + + return loss, {} + + +def edge_loss(points: torch.Tensor, gt_points: torch.Tensor) -> torch.Tensor: + device, dtype = points.device, points.dtype + height, width = points.shape[-3:-1] + + mask = torch.isfinite(gt_points).all(dim=-1) + gt_points = torch.where(mask[..., None], gt_points, 1) + + dx = points[..., :-1, :, :] - points[..., 1:, :, :] + dy = points[..., :, :-1, :] - points[..., :, 1:, :] + + gt_dx = gt_points[..., :-1, :, :] - gt_points[..., 1:, :, :] + gt_dy = gt_points[..., :, :-1, :] - gt_points[..., :, 1:, :] + + mask_dx = mask[..., :-1, :] & mask[..., 1:, :] + mask_dy = mask[..., :, :-1] & mask[..., :, 1:] + + MIN_ANGLE, MAX_ANGLE, BETA_RAD = math.radians(0.1), math.radians(90), math.radians(3) + + loss_dx = mask_dx * _smooth(angle_diff_vec3(dx, gt_dx).clamp(MIN_ANGLE, MAX_ANGLE), beta=BETA_RAD) + loss_dy = mask_dy * _smooth(angle_diff_vec3(dy, gt_dy).clamp(MIN_ANGLE, MAX_ANGLE), beta=BETA_RAD) + loss = (loss_dx.mean(dim=(-2, -1)) + loss_dy.mean(dim=(-2, -1))) / (2 * max(points.shape[-3:-1])) + + return loss, {} + + +def mask_l2_loss(pred_mask: torch.Tensor, gt_mask_pos: torch.Tensor, gt_mask_neg: torch.Tensor) -> torch.Tensor: + loss = gt_mask_neg.float() * pred_mask.square() + gt_mask_pos.float() * (1 - pred_mask).square() + loss = loss.mean(dim=(-2, -1)) + return loss, {} + + +def mask_bce_loss(pred_mask_prob: torch.Tensor, gt_mask_pos: torch.Tensor, gt_mask_neg: torch.Tensor) -> torch.Tensor: + loss = (gt_mask_pos | gt_mask_neg) * F.binary_cross_entropy(pred_mask_prob, gt_mask_pos.float(), reduction='none') + loss = loss.mean(dim=(-2, -1)) + return loss, {} + + +def metric_scale_loss(scale_pred: torch.Tensor, scale_gt: torch.Tensor): + valid = scale_gt > 0 + return torch.where(valid, F.mse_loss(scale_pred.log(), torch.where(valid, scale_gt.log(), 0), reduction='none'), 0), {} + + +def normal_map_loss(pred_normal: torch.Tensor, gt_normal: torch.Tensor) -> torch.Tensor: + mask = torch.isfinite(gt_normal).all(dim=-1) + gt_normal = torch.where(mask[..., None], gt_normal, 1) + + loss = (mask * utils3d.pt.angle_between(pred_normal, gt_normal).square()).mean(dim=(-2, -1)) + return loss, {} diff --git a/moge/train/utils.py b/moge/train/utils.py new file mode 100644 index 0000000..5f21e00 --- /dev/null +++ b/moge/train/utils.py @@ -0,0 +1,57 @@ +from typing import * +import fnmatch + +import sympy +import torch +import torch.nn as nn + + +def any_match(s: str, patterns: List[str]) -> bool: + return any(fnmatch.fnmatch(s, pat) for pat in patterns) + + +def build_optimizer(model: nn.Module, optimizer_config: Dict[str, Any]) -> torch.optim.Optimizer: + named_param_groups = [ + { + k: p for k, p in model.named_parameters() if any_match(k, param_group_config['params']['include']) and not any_match(k, param_group_config['params'].get('exclude', [])) + } for param_group_config in optimizer_config['params'] + ] + excluded_params = [k for k, p in model.named_parameters() if p.requires_grad and not any(k in named_params for named_params in named_param_groups)] + assert len(excluded_params) == 0, f'The following parameters require grad but are excluded from the optimizer: {excluded_params}' + optimizer_cls = getattr(torch.optim, optimizer_config['type']) + optimizer = optimizer_cls([ + { + **param_group_config, + 'params': list(params.values()), + } for param_group_config, params in zip(optimizer_config['params'], named_param_groups) + ]) + return optimizer + + +def parse_lr_lambda(s: str) -> Callable[[int], float]: + epoch = sympy.symbols('epoch') + lr_lambda = sympy.sympify(s) + return sympy.lambdify(epoch, lr_lambda, 'math') + + +def build_lr_scheduler(optimizer: torch.optim.Optimizer, scheduler_config: Dict[str, Any]) -> torch.optim.lr_scheduler._LRScheduler: + if scheduler_config['type'] == "SequentialLR": + child_schedulers = [ + build_lr_scheduler(optimizer, child_scheduler_config) + for child_scheduler_config in scheduler_config['params']['schedulers'] + ] + return torch.optim.lr_scheduler.SequentialLR(optimizer, schedulers=child_schedulers, milestones=scheduler_config['params']['milestones']) + elif scheduler_config['type'] == "LambdaLR": + lr_lambda = scheduler_config['params']['lr_lambda'] + if isinstance(lr_lambda, str): + lr_lambda = parse_lr_lambda(lr_lambda) + elif isinstance(lr_lambda, list): + lr_lambda = [parse_lr_lambda(l) for l in lr_lambda] + return torch.optim.lr_scheduler.LambdaLR( + optimizer, + lr_lambda=lr_lambda, + ) + else: + scheduler_cls = getattr(torch.optim.lr_scheduler, scheduler_config['type']) + scheduler = scheduler_cls(optimizer, **scheduler_config.get('params', {})) + return scheduler \ No newline at end of file diff --git a/moge/utils/__init__.py b/moge/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/moge/utils/alignment.py b/moge/utils/alignment.py new file mode 100644 index 0000000..3d6bb78 --- /dev/null +++ b/moge/utils/alignment.py @@ -0,0 +1,416 @@ +from typing import * +import math +from collections import namedtuple + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.types +import utils3d + + +def scatter_min(size: int, dim: int, index: torch.LongTensor, src: torch.Tensor) -> torch.return_types.min: + "Scatter the minimum value along the given dimension of `input` into `src` at the indices specified in `index`." + shape = src.shape[:dim] + (size,) + src.shape[dim + 1:] + minimum = torch.full(shape, float('inf'), dtype=src.dtype, device=src.device).scatter_reduce(dim=dim, index=index, src=src, reduce='amin', include_self=False) + minimum_where = torch.where(src == torch.gather(minimum, dim=dim, index=index)) + indices = torch.full(shape, -1, dtype=torch.long, device=src.device) + indices[(*minimum_where[:dim], index[minimum_where], *minimum_where[dim + 1:])] = minimum_where[dim] + return torch.return_types.min((minimum, indices)) + + +def split_batch_fwd(fn: Callable, chunk_size: int, *args, **kwargs): + batch_size = next(x for x in (*args, *kwargs.values()) if isinstance(x, torch.Tensor)).shape[0] + n_chunks = batch_size // chunk_size + (batch_size % chunk_size > 0) + splited_args = tuple(arg.split(chunk_size, dim=0) if isinstance(arg, torch.Tensor) else [arg] * n_chunks for arg in args) + splited_kwargs = {k: [v.split(chunk_size, dim=0) if isinstance(v, torch.Tensor) else [v] * n_chunks] for k, v in kwargs.items()} + results = [] + for i in range(n_chunks): + chunk_args = tuple(arg[i] for arg in splited_args) + chunk_kwargs = {k: v[i] for k, v in splited_kwargs.items()} + results.append(fn(*chunk_args, **chunk_kwargs)) + + if isinstance(results[0], tuple): + return tuple(torch.cat(r, dim=0) for r in zip(*results)) + else: + return torch.cat(results, dim=0) + + +def _pad_inf(x_: torch.Tensor): + return torch.cat([torch.full_like(x_[..., :1], -torch.inf), x_, torch.full_like(x_[..., :1], torch.inf)], dim=-1) + + +def _pad_cumsum(cumsum: torch.Tensor): + return torch.cat([torch.zeros_like(cumsum[..., :1]), cumsum, cumsum[..., -1:]], dim=-1) + + +def _compute_residual(a: torch.Tensor, xyw: torch.Tensor, trunc: float): + return a.mul(xyw[..., 0]).sub_(xyw[..., 1]).abs_().mul_(xyw[..., 2]).clamp_max_(trunc).sum(dim=-1) + + +def align(x: torch.Tensor, y: torch.Tensor, w: torch.Tensor, trunc: Optional[Union[float, torch.Tensor]] = None, eps: float = 1e-7) -> Tuple[torch.Tensor, torch.Tensor, torch.LongTensor]: + """ + If trunc is None, solve `min sum_i w_i * |a * x_i - y_i|`, otherwise solve `min sum_i min(trunc, w_i * |a * x_i - y_i|)`. + + w_i must be >= 0. + + ### Parameters: + - `x`: tensor of shape (..., n) + - `y`: tensor of shape (..., n) + - `w`: tensor of shape (..., n) + - `trunc`: optional, float or tensor of shape (..., n) or None + + ### Returns: + - `a`: tensor of shape (...), differentiable + - `loss`: tensor of shape (...), value of loss function at `a`, detached + - `index`: tensor of shape (...), where a = y[idx] / x[idx] + """ + if trunc is None: + x, y, w = torch.broadcast_tensors(x, y, w) + sign = torch.sign(x) + x, y = x * sign, y * sign + y_div_x = y / x.clamp_min(eps) + y_div_x, argsort = y_div_x.sort(dim=-1) + + wx = torch.gather(x * w, dim=-1, index=argsort) + derivatives = 2 * wx.cumsum(dim=-1) - wx.sum(dim=-1, keepdim=True) + search = torch.searchsorted(derivatives, torch.zeros_like(derivatives[..., :1]), side='left').clamp_max(derivatives.shape[-1] - 1) + + a = y_div_x.gather(dim=-1, index=search).squeeze(-1) + index = argsort.gather(dim=-1, index=search).squeeze(-1) + loss = (w * (a[..., None] * x - y).abs()).sum(dim=-1) + + else: + # Reshape to (batch_size, n) for simplicity + x, y, w = torch.broadcast_tensors(x, y, w) + batch_shape = x.shape[:-1] + batch_size = math.prod(batch_shape) + x, y, w = x.reshape(-1, x.shape[-1]), y.reshape(-1, y.shape[-1]), w.reshape(-1, w.shape[-1]) + + sign = torch.sign(x) + x, y = x * sign, y * sign + wx, wy = w * x, w * y + xyw = torch.stack([x, y, w], dim=-1) # Stacked for convenient gathering + + y_div_x = A = y / x.clamp_min(eps) + B = (wy - trunc) / wx.clamp_min(eps) + C = (wy + trunc) / wx.clamp_min(eps) + with torch.no_grad(): + # Caculate prefix sum by orders of A, B, C + A, A_argsort = A.sort(dim=-1) + Q_A = torch.cumsum(torch.gather(wx, dim=-1, index=A_argsort), dim=-1) + A, Q_A = _pad_inf(A), _pad_cumsum(Q_A) # Pad [-inf, A1, ..., An, inf] and [0, Q1, ..., Qn, Qn] to handle edge cases. + + B, B_argsort = B.sort(dim=-1) + Q_B = torch.cumsum(torch.gather(wx, dim=-1, index=B_argsort), dim=-1) + B, Q_B = _pad_inf(B), _pad_cumsum(Q_B) + + C, C_argsort = C.sort(dim=-1) + Q_C = torch.cumsum(torch.gather(wx, dim=-1, index=C_argsort), dim=-1) + C, Q_C = _pad_inf(C), _pad_cumsum(Q_C) + + # Caculate left and right derivative of A + j_A = torch.searchsorted(A, y_div_x, side='left').sub_(1) + j_B = torch.searchsorted(B, y_div_x, side='left').sub_(1) + j_C = torch.searchsorted(C, y_div_x, side='left').sub_(1) + left_derivative = 2 * torch.gather(Q_A, dim=-1, index=j_A) - torch.gather(Q_B, dim=-1, index=j_B) - torch.gather(Q_C, dim=-1, index=j_C) + j_A = torch.searchsorted(A, y_div_x, side='right').sub_(1) + j_B = torch.searchsorted(B, y_div_x, side='right').sub_(1) + j_C = torch.searchsorted(C, y_div_x, side='right').sub_(1) + right_derivative = 2 * torch.gather(Q_A, dim=-1, index=j_A) - torch.gather(Q_B, dim=-1, index=j_B) - torch.gather(Q_C, dim=-1, index=j_C) + + # Find extrema + is_extrema = (left_derivative < 0) & (right_derivative >= 0) + is_extrema[..., 0] |= ~is_extrema.any(dim=-1) # In case all derivatives are zero, take the first one as extrema. + where_extrema_batch, where_extrema_index = torch.where(is_extrema) + + # Calculate objective value at extrema + extrema_a = y_div_x[where_extrema_batch, where_extrema_index] # (num_extrema,) + MAX_ELEMENTS = 4096 ** 2 # Split into small batches to avoid OOM in case there are too many extrema.(~1G) + SPLIT_SIZE = MAX_ELEMENTS // x.shape[-1] + extrema_value = torch.cat([ + _compute_residual(extrema_a_split[:, None], xyw[extrema_i_split, :, :], trunc) + for extrema_a_split, extrema_i_split in zip(extrema_a.split(SPLIT_SIZE), where_extrema_batch.split(SPLIT_SIZE)) + ]) # (num_extrema,) + + # Find minima among corresponding extrema + minima, indices = scatter_min(size=batch_size, dim=0, index=where_extrema_batch, src=extrema_value) # (batch_size,) + index = where_extrema_index[indices] + + a = torch.gather(y, dim=-1, index=index[..., None]) / torch.gather(x, dim=-1, index=index[..., None]).clamp_min(eps) + a = a.reshape(batch_shape) + loss = minima.reshape(batch_shape) + index = index.reshape(batch_shape) + + return a, loss, index + + +def align_depth_scale(depth_src: torch.Tensor, depth_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None): + """ + Align `depth_src` to `depth_tgt` with given constant weights. + + ### Parameters: + - `depth_src: torch.Tensor` of shape (..., N) + - `depth_tgt: torch.Tensor` of shape (..., N) + + """ + scale, _, _ = align(depth_src, depth_tgt, weight, trunc) + + return scale + + +def align_depth_affine(depth_src: torch.Tensor, depth_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None): + """ + Align `depth_src` to `depth_tgt` with given constant weights. + + ### Parameters: + - `depth_src: torch.Tensor` of shape (..., N) + - `depth_tgt: torch.Tensor` of shape (..., N) + - `weight: torch.Tensor` of shape (..., N) + - `trunc: float` or tensor of shape (..., N) or None + + ### Returns: + - `scale: torch.Tensor` of shape (...). + - `shift: torch.Tensor` of shape (...). + """ + dtype, device = depth_src.dtype, depth_src.device + + # Flatten batch dimensions for simplicity + batch_shape, n = depth_src.shape[:-1], depth_src.shape[-1] + batch_size = math.prod(batch_shape) + depth_src, depth_tgt, weight = depth_src.reshape(batch_size, n), depth_tgt.reshape(batch_size, n), weight.reshape(batch_size, n) + + # Here, we take anchors only for non-zero weights. + # Although the results will be still correct even anchor points have zero weight, + # it is wasting computation and may cause instability in some cases, e.g. too many extrema. + anchors_where_batch, anchors_where_n = torch.where(weight > 0) + + # Stop gradient when solving optimal anchors + with torch.no_grad(): + depth_src_anchor = depth_src[anchors_where_batch, anchors_where_n] # (anchors) + depth_tgt_anchor = depth_tgt[anchors_where_batch, anchors_where_n] # (anchors) + + depth_src_anchored = depth_src[anchors_where_batch, :] - depth_src_anchor[..., None] # (anchors, n) + depth_tgt_anchored = depth_tgt[anchors_where_batch, :] - depth_tgt_anchor[..., None] # (anchors, n) + weight_anchored = weight[anchors_where_batch, :] # (anchors, n) + + scale, loss, index = align(depth_src_anchored, depth_tgt_anchored, weight_anchored, trunc) # (anchors) + + loss, index_anchor = scatter_min(size=batch_size, dim=0, index=anchors_where_batch, src=loss) # (batch_size,) + + # Reproduce by indexing for shorter compute graph + index_1 = anchors_where_n[index_anchor] # (batch_size,) + index_2 = index[index_anchor] # (batch_size,) + + tgt_1, src_1 = torch.gather(depth_tgt, dim=1, index=index_1[..., None]).squeeze(-1), torch.gather(depth_src, dim=1, index=index_1[..., None]).squeeze(-1) + tgt_2, src_2 = torch.gather(depth_tgt, dim=1, index=index_2[..., None]).squeeze(-1), torch.gather(depth_src, dim=1, index=index_2[..., None]).squeeze(-1) + + scale = (tgt_2 - tgt_1) / torch.where(src_2 != src_1, src_2 - src_1, 1e-7) + shift = tgt_1 - scale * src_1 + + scale, shift = scale.reshape(batch_shape), shift.reshape(batch_shape) + + return scale, shift + +def align_depth_affine_irls(depth_src: torch.Tensor, depth_tgt: torch.Tensor, weight: Optional[torch.Tensor], max_iter: int = 100, eps: float = 1e-12): + """ + Align `depth_src` to `depth_tgt` with given constant weights using IRLS. + """ + dtype, device = depth_src.dtype, depth_src.device + + w = weight + x = torch.stack([depth_src, torch.ones_like(depth_src)], dim=-1) + y = depth_tgt + + for i in range(max_iter): + beta = (x.transpose(-1, -2) @ (w * y)) @ (x.transpose(-1, -2) @ (w[..., None] * x)).inverse().transpose(-2, -1) + w = 1 / (y - (x @ beta[..., None])[..., 0]).abs().clamp_min(eps) + + return beta[..., 0], beta[..., 1] + + +def align_points_scale(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None): + """ + ### Parameters: + - `points_src: torch.Tensor` of shape (..., N, 3) + - `points_tgt: torch.Tensor` of shape (..., N, 3) + - `weight: torch.Tensor` of shape (..., N) + + ### Returns: + - `a: torch.Tensor` of shape (...). Only positive solutions are garunteed. You should filter out negative scales before using it. + - `b: torch.Tensor` of shape (...) + """ + dtype, device = points_src.dtype, points_src.device + + scale, _, _ = align(points_src.flatten(-2), points_tgt.flatten(-2), weight[..., None].expand_as(points_src).flatten(-2), trunc) + + return scale + + +def align_points_scale_z_shift(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None): + """ + Align `points_src` to `points_tgt` with respect to a shared xyz scale and z shift. + It is similar to `align_affine` but scale and shift are applied to different dimensions. + + ### Parameters: + - `points_src: torch.Tensor` of shape (..., N, 3) + - `points_tgt: torch.Tensor` of shape (..., N, 3) + - `weights: torch.Tensor` of shape (..., N) + + ### Returns: + - `scale: torch.Tensor` of shape (...). + - `shift: torch.Tensor` of shape (..., 3). x and y shifts are zeros. + """ + dtype, device = points_src.dtype, points_src.device + + # Flatten batch dimensions for simplicity + batch_shape, n = points_src.shape[:-2], points_src.shape[-2] + batch_size = math.prod(batch_shape) + points_src, points_tgt, weight = points_src.reshape(batch_size, n, 3), points_tgt.reshape(batch_size, n, 3), weight.reshape(batch_size, n) + + # Take anchors + anchor_where_batch, anchor_where_n = torch.where(weight > 0) + with torch.no_grad(): + zeros = torch.zeros(anchor_where_batch.shape[0], device=device, dtype=dtype) + points_src_anchor = torch.stack([zeros, zeros, points_src[anchor_where_batch, anchor_where_n, 2]], dim=-1) # (anchors, 3) + points_tgt_anchor = torch.stack([zeros, zeros, points_tgt[anchor_where_batch, anchor_where_n, 2]], dim=-1) # (anchors, 3) + + points_src_anchored = points_src[anchor_where_batch, :, :] - points_src_anchor[..., None, :] # (anchors, n, 3) + points_tgt_anchored = points_tgt[anchor_where_batch, :, :] - points_tgt_anchor[..., None, :] # (anchors, n, 3) + weight_anchored = weight[anchor_where_batch, :, None].expand(-1, -1, 3) # (anchors, n, 3) + + # Solve optimal scale and shift for each anchor + MAX_ELEMENTS = 2 ** 20 + scale, loss, index = split_batch_fwd(align, MAX_ELEMENTS // n, points_src_anchored.flatten(-2), points_tgt_anchored.flatten(-2), weight_anchored.flatten(-2), trunc) # (anchors,) + + loss, index_anchor = scatter_min(size=batch_size, dim=0, index=anchor_where_batch, src=loss) # (batch_size,) + + # Reproduce by indexing for shorter compute graph + index_2 = index[index_anchor] # (batch_size,) [0, 3n) + index_1 = anchor_where_n[index_anchor] * 3 + index_2 % 3 # (batch_size,) [0, 3n) + + zeros = torch.zeros((batch_size, n), device=device, dtype=dtype) + points_tgt_00z, points_src_00z = torch.stack([zeros, zeros, points_tgt[..., 2]], dim=-1), torch.stack([zeros, zeros, points_src[..., 2]], dim=-1) + tgt_1, src_1 = torch.gather(points_tgt_00z.flatten(-2), dim=1, index=index_1[..., None]).squeeze(-1), torch.gather(points_src_00z.flatten(-2), dim=1, index=index_1[..., None]).squeeze(-1) + tgt_2, src_2 = torch.gather(points_tgt.flatten(-2), dim=1, index=index_2[..., None]).squeeze(-1), torch.gather(points_src.flatten(-2), dim=1, index=index_2[..., None]).squeeze(-1) + + scale = (tgt_2 - tgt_1) / torch.where(src_2 != src_1, src_2 - src_1, 1.0) + shift = torch.gather(points_tgt_00z, dim=1, index=(index_1 // 3)[..., None, None].expand(-1, -1, 3)).squeeze(-2) - scale[..., None] * torch.gather(points_src_00z, dim=1, index=(index_1 // 3)[..., None, None].expand(-1, -1, 3)).squeeze(-2) + scale, shift = scale.reshape(batch_shape), shift.reshape(*batch_shape, 3) + + return scale, shift + + +def align_points_scale_xyz_shift(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None, max_iters: int = 30, eps: float = 1e-6): + """ + Align `points_src` to `points_tgt` with respect to a shared xyz scale and z shift. + It is similar to `align_affine` but scale and shift are applied to different dimensions. + + ### Parameters: + - `points_src: torch.Tensor` of shape (..., N, 3) + - `points_tgt: torch.Tensor` of shape (..., N, 3) + - `weights: torch.Tensor` of shape (..., N) + + ### Returns: + - `scale: torch.Tensor` of shape (...). + - `shift: torch.Tensor` of shape (..., 3) + """ + dtype, device = points_src.dtype, points_src.device + + # Flatten batch dimensions for simplicity + batch_shape, n = points_src.shape[:-2], points_src.shape[-2] + batch_size = math.prod(batch_shape) + points_src, points_tgt, weight = points_src.reshape(batch_size, n, 3), points_tgt.reshape(batch_size, n, 3), weight.reshape(batch_size, n) + + # Take anchors + anchor_where_batch, anchor_where_n = torch.where(weight > 0) + + with torch.no_grad(): + points_src_anchor = points_src[anchor_where_batch, anchor_where_n] # (anchors, 3) + points_tgt_anchor = points_tgt[anchor_where_batch, anchor_where_n] # (anchors, 3) + + points_src_anchored = points_src[anchor_where_batch, :, :] - points_src_anchor[..., None, :] # (anchors, n, 3) + points_tgt_anchored = points_tgt[anchor_where_batch, :, :] - points_tgt_anchor[..., None, :] # (anchors, n, 3) + weight_anchored = weight[anchor_where_batch, :, None].expand(-1, -1, 3) # (anchors, n, 3) + + # Solve optimal scale and shift for each anchor + MAX_ELEMENTS = 2 ** 20 + scale, loss, index = split_batch_fwd(align, MAX_ELEMENTS // 2, points_src_anchored.flatten(-2), points_tgt_anchored.flatten(-2), weight_anchored.flatten(-2), trunc) # (anchors,) + + # Get optimal scale and shift for each batch element + loss, index_anchor = scatter_min(size=batch_size, dim=0, index=anchor_where_batch, src=loss) # (batch_size,) + + index_2 = index[index_anchor] # (batch_size,) [0, 3n) + index_1 = anchor_where_n[index_anchor] * 3 + index_2 % 3 # (batch_size,) [0, 3n) + + src_1, tgt_1 = torch.gather(points_src.flatten(-2), dim=1, index=index_1[..., None]).squeeze(-1), torch.gather(points_tgt.flatten(-2), dim=1, index=index_1[..., None]).squeeze(-1) + src_2, tgt_2 = torch.gather(points_src.flatten(-2), dim=1, index=index_2[..., None]).squeeze(-1), torch.gather(points_tgt.flatten(-2), dim=1, index=index_2[..., None]).squeeze(-1) + + scale = (tgt_2 - tgt_1) / torch.where(src_2 != src_1, src_2 - src_1, 1.0) + shift = torch.gather(points_tgt, dim=1, index=(index_1 // 3)[..., None, None].expand(-1, -1, 3)).squeeze(-2) - scale[..., None] * torch.gather(points_src, dim=1, index=(index_1 // 3)[..., None, None].expand(-1, -1, 3)).squeeze(-2) + + scale, shift = scale.reshape(batch_shape), shift.reshape(*batch_shape, 3) + + return scale, shift + + +def align_points_z_shift(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None, max_iters: int = 30, eps: float = 1e-6): + """ + Align `points_src` to `points_tgt` with respect to a Z-axis shift. + + ### Parameters: + - `points_src: torch.Tensor` of shape (..., N, 3) + - `points_tgt: torch.Tensor` of shape (..., N, 3) + - `weights: torch.Tensor` of shape (..., N) + + ### Returns: + - `scale: torch.Tensor` of shape (...). + - `shift: torch.Tensor` of shape (..., 3) + """ + dtype, device = points_src.dtype, points_src.device + + shift, _, _ = align(torch.ones_like(points_src[..., 2]), points_tgt[..., 2] - points_src[..., 2], weight, trunc) + shift = torch.stack([torch.zeros_like(shift), torch.zeros_like(shift), shift], dim=-1) + + return shift + + +def align_points_xyz_shift(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None, max_iters: int = 30, eps: float = 1e-6): + """ + Align `points_src` to `points_tgt` with respect to a Z-axis shift. + + ### Parameters: + - `points_src: torch.Tensor` of shape (..., N, 3) + - `points_tgt: torch.Tensor` of shape (..., N, 3) + - `weights: torch.Tensor` of shape (..., N) + + ### Returns: + - `scale: torch.Tensor` of shape (...). + - `shift: torch.Tensor` of shape (..., 3) + """ + dtype, device = points_src.dtype, points_src.device + + shift, _, _ = align(torch.ones_like(points_src).swapaxes(-2, -1), (points_tgt - points_src).swapaxes(-2, -1), weight[..., None, :], trunc) + + return shift + + +def align_affine_lstsq(x: torch.Tensor, y: torch.Tensor, w: torch.Tensor = None) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Solve `min sum_i w_i * (a * x_i + b - y_i ) ^ 2`, where `a` and `b` are scalars, with respect to `a` and `b` using least squares. + + ### Parameters: + - `x: torch.Tensor` of shape (..., N) + - `y: torch.Tensor` of shape (..., N) + - `w: torch.Tensor` of shape (..., N) + + ### Returns: + - `a: torch.Tensor` of shape (...,) + - `b: torch.Tensor` of shape (...,) + """ + w_sqrt = torch.ones_like(x) if w is None else w.sqrt() + A = torch.stack([w_sqrt * x, torch.ones_like(x)], dim=-1) + B = (w_sqrt * y)[..., None] + a, b = torch.linalg.lstsq(A, B)[0].squeeze(-1).unbind(-1) + return a, b \ No newline at end of file diff --git a/moge/utils/data_augmentation.py b/moge/utils/data_augmentation.py new file mode 100644 index 0000000..9fc4c9d --- /dev/null +++ b/moge/utils/data_augmentation.py @@ -0,0 +1,250 @@ +import os +import json +import time +import random +from typing import * +import itertools +from numbers import Number +import io + +import numpy as np +import cv2 +from PIL import Image +import torch +import torchvision.transforms.v2.functional as TF +import utils3d +from scipy.signal import fftconvolve + +from ..utils.geometry_numpy import harmonic_mean_numpy, norm3d, depth_occlusion_edge_numpy + + +def sample_perspective( + src_intrinsics: np.ndarray, + tgt_aspect: float, + center_augmentation: float, + fov_range_absolute: Tuple[float, float], + fov_range_relative: Tuple[float, float], + rng: np.random.Generator = None +) -> Tuple[np.ndarray, np.ndarray]: + raw_horizontal, raw_vertical = abs(1.0 / src_intrinsics[0, 0]), abs(1.0 / src_intrinsics[1, 1]) + raw_fov_x, raw_fov_y = utils3d.np.intrinsics_to_fov(src_intrinsics) + + # 1. set target fov + fov_range_absolute_min, fov_range_absolute_max = fov_range_absolute + fov_range_relative_min, fov_range_relative_max = fov_range_relative + tgt_fov_x_min = min(fov_range_relative_min * raw_fov_x, utils3d.focal_to_fov(utils3d.fov_to_focal(fov_range_relative_min * raw_fov_y) / tgt_aspect)) + tgt_fov_x_max = min(fov_range_relative_max * raw_fov_x, utils3d.focal_to_fov(utils3d.fov_to_focal(fov_range_relative_max * raw_fov_y) / tgt_aspect)) + tgt_fov_x_min, tgt_fov_max = max(np.deg2rad(fov_range_absolute_min), tgt_fov_x_min), min(np.deg2rad(fov_range_absolute_max), tgt_fov_x_max) + tgt_fov_x = rng.uniform(min(tgt_fov_x_min, tgt_fov_x_max), tgt_fov_x_max) + tgt_fov_y = utils3d.focal_to_fov(utils3d.np.fov_to_focal(tgt_fov_x) * tgt_aspect) + + # 2. set target image center (principal point) and the corresponding z-direction in raw camera space + center_dtheta = center_augmentation * rng.uniform(-0.5, 0.5) * (raw_fov_x - tgt_fov_x) + center_dphi = center_augmentation * rng.uniform(-0.5, 0.5) * (raw_fov_y - tgt_fov_y) + cu, cv = 0.5 + 0.5 * np.tan(center_dtheta) / np.tan(raw_fov_x / 2), 0.5 + 0.5 * np.tan(center_dphi) / np.tan(raw_fov_y / 2) + direction = utils3d.np.unproject_cv(np.array([[cu, cv]], dtype=np.float32), np.array([1.0], dtype=np.float32), intrinsics=src_intrinsics)[0] + + # 3. obtain the rotation matrix for homography warping (new_ext = R * old_ext) + R = utils3d.np.rotation_matrix_from_vectors(direction, np.array([0, 0, 1], dtype=np.float32)) + + # 4. shrink the target view to fit into the warped image + corners = np.array([[0, 0], [0, 1], [1, 1], [1, 0]], dtype=np.float32) + corners = np.concatenate([corners, np.ones((4, 1), dtype=np.float32)], axis=1) @ (np.linalg.inv(src_intrinsics).T @ R.T) # corners in viewport's camera plane + corners = corners[:, :2] / corners[:, 2:3] + tgt_horizontal, tgt_vertical = np.tan(tgt_fov_x / 2) * 2, np.tan(tgt_fov_y / 2) * 2 + warp_horizontal, warp_vertical = float('inf'), float('inf') + for i in range(4): + intersection, _ = utils3d.np.ray_intersection( + np.array([0., 0.]), np.array([[tgt_aspect, 1.0], [tgt_aspect, -1.0]]), + corners[i - 1], corners[i] - corners[i - 1], + ) + warp_horizontal, warp_vertical = min(warp_horizontal, 2 * np.abs(intersection[:, 0]).min()), min(warp_vertical, 2 * np.abs(intersection[:, 1]).min()) + tgt_horizontal, tgt_vertical = min(tgt_horizontal, warp_horizontal), min(tgt_vertical, warp_vertical) + + # 5. obtain the target intrinsics + fx, fy = 1 / tgt_horizontal, 1 / tgt_vertical + tgt_intrinsics = utils3d.np.intrinsics_from_focal_center(fx, fy, 0.5, 0.5).astype(np.float32) + + return tgt_intrinsics, R + + +def warp_perspective( + src_map: np.ndarray = None, + transform: np.ndarray = None, + tgt_size: Tuple[int, int] = None, + interpolation: Literal['nearest', 'bilinear', 'lanczos'] = 'nearest', + sparse_mask: np.ndarray = None, +): + """Perspective warping with careful resampling. + - For `lanczos`, use PIL to resize first to reduce aliasing. + - For `nearest` with sparse input, use mask-aware nearest resize to avoid losing points. + - For `bilinear` or `nearest` with dense input, directly use cv2.remap. + + - `transform` is the matrix that transforms homogeneous pixel coordinates of source image to those of target image, i.e., `p_tgt = transform @ p_src`. + """ + + tgt_height, tgt_width = tgt_size + src_height, src_width = src_map.shape[:2] + + # source to target transform + transform_pixel = np.array([[tgt_width, 0, -0.5], [0, tgt_height, -0.5], [0, 0, 1]], dtype=np.float32) @ transform @ np.array([[1 / src_width, 0, 0.5 / src_width], [0, 1 / src_height, 0.5 / src_height], [0, 0, 1]], dtype=np.float32) + # Get scale factor at the target center + w = np.dot(np.linalg.inv(transform_pixel)[2, :], np.array([tgt_width / 2, tgt_height / 2, 1], dtype=np.float32)) + scale_x, scale_y = w * np.linalg.norm(transform_pixel[:2, :2], axis=0) + + if interpolation == 'lanczos' and (scale_x < 0.8 or scale_y < 0.8): + # If lanczos & downsampling, use PIL to resize first to reduce aliasing + src_height, src_width = max(round(src_height * scale_y * 1.25), 16), max(round(src_width * scale_x * 1.25), 16) + src_map = np.array(Image.fromarray(src_map).resize((src_width, src_height), Image.Resampling.LANCZOS)) + elif interpolation == 'nearest' and sparse_mask is not None and (scale_x < 1 or scale_y < 1): + # If nearest and sparse, use mask-aware nearest resize first to avoid losing points + src_height, src_width = max(round(src_height * scale_y), 16), max(round(src_width * scale_x), 16) + src_map, _ = utils3d.np.masked_nearest_resize(src_map, mask=sparse_mask, size=(src_height, src_width)) + + # Recompute the pixel-space transform after resizing + transform_pixel = np.array([[tgt_width, 0, -0.5], [0, tgt_height, -0.5], [0, 0, 1]], dtype=np.float32) @ transform @ np.array([[1 / src_width, 0, 0.5 / src_width], [0, 1 / src_height, 0.5 / src_height], [0, 0, 1]], dtype=np.float32) + + # Remap + cv2_interpolation = {'nearest': cv2.INTER_NEAREST, 'bilinear': cv2.INTER_LINEAR, 'lanczos': cv2.INTER_LANCZOS4}[interpolation] + tgt_map = cv2.warpPerspective(src_map, transform_pixel, (tgt_width, tgt_height), flags=cv2_interpolation) + + return tgt_map + + +def image_color_augmentation(image: np.ndarray, augmentations: List[Dict[str, Any]], rng: np.random.Generator = None, depth: np.ndarray = None): + height, width = image.shape[:2] + if rng is None: + rng = np.random.default_rng() + if 'jittering' in augmentations: + image = torch.from_numpy(image).permute(2, 0, 1) + image = TF.adjust_brightness(image, rng.uniform(0.9, 1.1)) + image = TF.adjust_contrast(image, rng.uniform(0.9, 1.1)) + image = TF.adjust_saturation(image, rng.uniform(0.9, 1.1)) + image = TF.adjust_hue(image, rng.uniform(-0.05, 0.05)) + image = TF.adjust_gamma(image, rng.uniform(0.9, 1.1)) + image = image.permute(1, 2, 0).numpy() + if 'dof' in augmentations: + assert depth is not None, 'Depth map is required for DOF augmentation' + if rng.uniform() < 0.5: + dof_strength = rng.integers(12) + disp = 1 / depth + finite_mask = np.isfinite(depth) + disp_min, disp_max = disp[finite_mask].min(), disp[finite_mask].max() + disp = cv2.inpaint(np.nan_to_num(disp, nan=1), np.isnan(disp).astype(np.uint8), 3, cv2.INPAINT_TELEA).clip(0, disp_max) + dof_focus = rng.uniform(disp_min, disp_max) + image = depth_of_field(image, disp, dof_focus, dof_strength) + if 'shot_noise' in augmentations: + if rng.uniform() < 0.5: + k = np.exp(rng.uniform(np.log(100), np.log(10000))) / 255 + image = (rng.poisson(image * k) / k).clip(0, 255).astype(np.uint8) + if 'blurring' in augmentations: + if rng.uniform() < 0.5: + ratio = rng.uniform(0.25, 1) + image = cv2.resize(cv2.resize(image, (int(width * ratio), int(height * ratio)), interpolation=cv2.INTER_AREA), (width, height), interpolation=rng.choice([cv2.INTER_LINEAR_EXACT, cv2.INTER_CUBIC, cv2.INTER_LANCZOS4])) + if 'jpeg_loss' in augmentations: + if rng.uniform() < 0.5: + image = cv2.imdecode(cv2.imencode('.jpg', image, [cv2.IMWRITE_JPEG_QUALITY, rng.integers(20, 100)])[1], cv2.IMREAD_COLOR) + + return image + + + +def disk_kernel(radius: int) -> np.ndarray: + """ + Generate disk kernel with given radius. + + Args: + radius (int): Radius of the disk (in pixels). + + Returns: + np.ndarray: (2*radius+1, 2*radius+1) normalized convolution kernel. + """ + # Create coordinate grid centered at (0,0) + L = np.arange(-radius, radius + 1) + X, Y = np.meshgrid(L, L) + # Generate disk: region inside circle with radius R is 1 + kernel = ((X**2 + Y**2) <= radius**2).astype(np.float32) + # Normalize the kernel + kernel /= np.sum(kernel) + return kernel + + +def disk_blur(image: np.ndarray, radius: int) -> np.ndarray: + """ + Apply disk blur to an image using FFT convolution. + + Args: + image (np.ndarray): Input image, can be grayscale or color. + radius (int): Blur radius (in pixels). + + Returns: + np.ndarray: Blurred image. + """ + if radius == 0: + return image + kernel = disk_kernel(radius) + if image.ndim == 2: + blurred = fftconvolve(image, kernel, mode='same') + elif image.ndim == 3: + channels = [] + for i in range(image.shape[2]): + blurred_channel = fftconvolve(image[..., i], kernel, mode='same') + channels.append(blurred_channel) + blurred = np.stack(channels, axis=-1) + else: + raise ValueError("Image must be 2D or 3D.") + return blurred + + +def depth_of_field( + img: np.ndarray, + disp: np.ndarray, + focus_disp : float, + max_blur_radius : int = 10, +) -> np.ndarray: + """ + Apply depth of field effect to an image. + + Args: + img (numpy.ndarray): (H, W, 3) input image. + depth (numpy.ndarray): (H, W) depth map of the scene. + focus_depth (float): Focus depth of the lens. + strength (float): Strength of the depth of field effect. + max_blur_radius (int): Maximum blur radius (in pixels). + + Returns: + numpy.ndarray: (H, W, 3) output image with depth of field effect applied. + """ + # Precalculate dialated depth map for each blur radius + max_disp = np.max(disp) + disp = disp / max_disp + focus_disp = focus_disp / max_disp + dilated_disp = [] + for radius in range(max_blur_radius + 1): + dilated_disp.append(cv2.dilate(disp, cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2 * radius + 1, 2 * radius + 1)), iterations=1)) + + # Determine the blur radius for each pixel based on the depth map + blur_radii = np.clip(np.abs(disp - focus_disp) * max_blur_radius, 0, max_blur_radius).astype(np.int32) + for radius in range(max_blur_radius + 1): + dialted_blur_radii = np.clip(np.abs(dilated_disp[radius] - focus_disp) * max_blur_radius, 0, max_blur_radius).astype(np.int32) + mask = (dialted_blur_radii >= radius) & (dialted_blur_radii >= blur_radii) & (dilated_disp[radius] > disp) + blur_radii[mask] = dialted_blur_radii[mask] + blur_radii = np.clip(blur_radii, 0, max_blur_radius) + blur_radii = cv2.blur(blur_radii, (5, 5)) + + # Precalculate the blured image for each blur radius + unique_radii = np.unique(blur_radii) + precomputed = {} + for radius in range(max_blur_radius + 1): + if radius not in unique_radii: + continue + precomputed[radius] = disk_blur(img, radius) + + # Composit the blured image for each pixel + output = np.zeros_like(img) + for r in unique_radii: + mask = blur_radii == r + output[mask] = precomputed[r][mask] + + return output + diff --git a/moge/utils/download.py b/moge/utils/download.py new file mode 100644 index 0000000..886edbc --- /dev/null +++ b/moge/utils/download.py @@ -0,0 +1,55 @@ +from pathlib import Path +from typing import * +import requests + +from tqdm import tqdm + + +__all__ = ["download_file", "download_bytes"] + + +def download_file(url: str, filepath: Union[str, Path], headers: dict = None, resume: bool = True) -> None: + # Ensure headers is a dict if not provided + headers = headers or {} + + # Initialize local variables + file_path = Path(filepath) + downloaded_bytes = 0 + + # Check if we should resume the download + if resume and file_path.exists(): + downloaded_bytes = file_path.stat().st_size + headers['Range'] = f"bytes={downloaded_bytes}-" + + # Make a GET request to fetch the file + with requests.get(url, stream=True, headers=headers) as response: + response.raise_for_status() # This will raise an HTTPError if the status is 4xx/5xx + + # Calculate the total size to download + total_size = downloaded_bytes + int(response.headers.get('content-length', 0)) + + # Display a progress bar while downloading + with ( + tqdm(desc=f"Downloading {file_path.name}", total=total_size, unit='B', unit_scale=True, leave=False) as pbar, + open(file_path, 'ab') as file, + ): + # Set the initial position of the progress bar + pbar.update(downloaded_bytes) + + # Write the content to the file in chunks + for chunk in response.iter_content(chunk_size=4096): + file.write(chunk) + pbar.update(len(chunk)) + + +def download_bytes(url: str, headers: dict = None) -> bytes: + # Ensure headers is a dict if not provided + headers = headers or {} + + # Make a GET request to fetch the file + with requests.get(url, stream=True, headers=headers) as response: + response.raise_for_status() # This will raise an HTTPError if the status is 4xx/5xx + + # Read the content of the response + return response.content + \ No newline at end of file diff --git a/moge/utils/geometry_numpy.py b/moge/utils/geometry_numpy.py new file mode 100644 index 0000000..99de45c --- /dev/null +++ b/moge/utils/geometry_numpy.py @@ -0,0 +1,261 @@ +from typing import * +from functools import partial +import math + +import cv2 +import numpy as np +from scipy.signal import fftconvolve +import numpy as np +import utils3d + +from .tools import timeit + + +def weighted_mean_numpy(x: np.ndarray, w: np.ndarray = None, axis: Union[int, Tuple[int,...]] = None, keepdims: bool = False, eps: float = 1e-7) -> np.ndarray: + if w is None: + return np.mean(x, axis=axis) + else: + w = w.astype(x.dtype) + return (x * w).mean(axis=axis) / np.clip(w.mean(axis=axis), eps, None) + + +def harmonic_mean_numpy(x: np.ndarray, w: np.ndarray = None, axis: Union[int, Tuple[int,...]] = None, keepdims: bool = False, eps: float = 1e-7) -> np.ndarray: + if w is None: + return 1 / (1 / np.clip(x, eps, None)).mean(axis=axis) + else: + w = w.astype(x.dtype) + return 1 / (weighted_mean_numpy(1 / (x + eps), w, axis=axis, keepdims=keepdims, eps=eps) + eps) + + +def normalized_view_plane_uv_numpy(width: int, height: int, aspect_ratio: float = None, dtype: np.dtype = np.float32) -> np.ndarray: + "UV with left-top corner as (-width / diagonal, -height / diagonal) and right-bottom corner as (width / diagonal, height / diagonal)" + if aspect_ratio is None: + aspect_ratio = width / height + + span_x = aspect_ratio / (1 + aspect_ratio ** 2) ** 0.5 + span_y = 1 / (1 + aspect_ratio ** 2) ** 0.5 + + u = np.linspace(-span_x * (width - 1) / width, span_x * (width - 1) / width, width, dtype=dtype) + v = np.linspace(-span_y * (height - 1) / height, span_y * (height - 1) / height, height, dtype=dtype) + u, v = np.meshgrid(u, v, indexing='xy') + uv = np.stack([u, v], axis=-1) + return uv + + +def focal_to_fov_numpy(focal: np.ndarray): + return 2 * np.arctan(0.5 / focal) + + +def fov_to_focal_numpy(fov: np.ndarray): + return 0.5 / np.tan(fov / 2) + + +def intrinsics_to_fov_numpy(intrinsics: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + fov_x = focal_to_fov_numpy(intrinsics[..., 0, 0]) + fov_y = focal_to_fov_numpy(intrinsics[..., 1, 1]) + return fov_x, fov_y + + +def point_map_to_depth_legacy_numpy(points: np.ndarray): + height, width = points.shape[-3:-1] + diagonal = (height ** 2 + width ** 2) ** 0.5 + uv = normalized_view_plane_uv_numpy(width, height, dtype=points.dtype) # (H, W, 2) + _, uv = np.broadcast_arrays(points[..., :2], uv) + + # Solve least squares problem + b = (uv * points[..., 2:]).reshape(*points.shape[:-3], -1) # (..., H * W * 2) + A = np.stack([points[..., :2], -uv], axis=-1).reshape(*points.shape[:-3], -1, 2) # (..., H * W * 2, 2) + + M = A.swapaxes(-2, -1) @ A + solution = (np.linalg.inv(M + 1e-6 * np.eye(2)) @ (A.swapaxes(-2, -1) @ b[..., None])).squeeze(-1) + focal, shift = solution + + depth = points[..., 2] + shift[..., None, None] + fov_x = np.arctan(width / diagonal / focal) * 2 + fov_y = np.arctan(height / diagonal / focal) * 2 + return depth, fov_x, fov_y, shift + + +def solve_optimal_focal_shift(uv: np.ndarray, xyz: np.ndarray): + "Solve `min |focal * xy / (z + shift) - uv|` with respect to shift and focal" + from scipy.optimize import least_squares + uv, xy, z = uv.reshape(-1, 2), xyz[..., :2].reshape(-1, 2), xyz[..., 2].reshape(-1) + + def fn(uv: np.ndarray, xy: np.ndarray, z: np.ndarray, shift: np.ndarray): + xy_proj = xy / (z + shift)[: , None] + f = (xy_proj * uv).sum() / np.square(xy_proj).sum() + err = (f * xy_proj - uv).ravel() + return err + + solution = least_squares(partial(fn, uv, xy, z), x0=0, ftol=1e-3, method='lm') + optim_shift = solution['x'].squeeze().astype(np.float32) + + xy_proj = xy / (z + optim_shift)[: , None] + optim_focal = (xy_proj * uv).sum() / np.square(xy_proj).sum() + + return optim_shift, optim_focal + + +def solve_optimal_shift(uv: np.ndarray, xyz: np.ndarray, focal: float): + "Solve `min |focal * xy / (z + shift) - uv|` with respect to shift" + from scipy.optimize import least_squares + uv, xy, z = uv.reshape(-1, 2), xyz[..., :2].reshape(-1, 2), xyz[..., 2].reshape(-1) + + def fn(uv: np.ndarray, xy: np.ndarray, z: np.ndarray, shift: np.ndarray): + xy_proj = xy / (z + shift)[: , None] + err = (focal * xy_proj - uv).ravel() + return err + + solution = least_squares(partial(fn, uv, xy, z), x0=0, ftol=1e-3, method='lm') + optim_shift = solution['x'].squeeze().astype(np.float32) + + return optim_shift + + +def recover_focal_shift_numpy(points: np.ndarray, mask: np.ndarray = None, focal: float = None, downsample_size: Tuple[int, int] = (64, 64)): + import cv2 + assert points.shape[-1] == 3, "Points should (H, W, 3)" + + height, width = points.shape[-3], points.shape[-2] + diagonal = (height ** 2 + width ** 2) ** 0.5 + + uv = normalized_view_plane_uv_numpy(width=width, height=height) + + if mask is None: + points_lr = cv2.resize(points, downsample_size, interpolation=cv2.INTER_LINEAR).reshape(-1, 3) + uv_lr = cv2.resize(uv, downsample_size, interpolation=cv2.INTER_LINEAR).reshape(-1, 2) + else: + points_lr, uv_lr, mask_lr = utils3d.np.masked_nearest_resize(points, uv, mask=mask, size=downsample_size) + + if points_lr.size < 2: + return 1., 0. + + if focal is None: + shift,focal = solve_optimal_focal_shift(uv_lr, points_lr) + else: + shift = solve_optimal_shift(uv_lr, points_lr, focal) + + return focal, shift + + +def norm3d(x: np.ndarray) -> np.ndarray: + "Faster `np.linalg.norm(x, axis=-1)` for 3D vectors" + return np.sqrt(np.square(x[..., 0]) + np.square(x[..., 1]) + np.square(x[..., 2])) + + +def depth_occlusion_edge_numpy(depth: np.ndarray, mask: np.ndarray, thickness: int = 1, tol: float = 0.1): + disp = np.where(mask, 1 / depth, 0) + disp_pad = np.pad(disp, (thickness, thickness), constant_values=0) + mask_pad = np.pad(mask, (thickness, thickness), constant_values=False) + kernel_size = 2 * thickness + 1 + disp_window = utils3d.np.sliding_window(disp_pad, (kernel_size, kernel_size), 1, axis=(-2, -1)) # [..., H, W, kernel_size ** 2] + mask_window = utils3d.np.sliding_window(mask_pad, (kernel_size, kernel_size), 1, axis=(-2, -1)) # [..., H, W, kernel_size ** 2] + + disp_mean = weighted_mean_numpy(disp_window, mask_window, axis=(-2, -1)) + fg_edge_mask = mask & (disp > (1 + tol) * disp_mean) + bg_edge_mask = mask & (disp_mean > (1 + tol) * disp) + + edge_mask = (cv2.dilate(fg_edge_mask.astype(np.uint8), np.ones((3, 3), dtype=np.uint8), iterations=thickness) > 0) \ + & (cv2.dilate(bg_edge_mask.astype(np.uint8), np.ones((3, 3), dtype=np.uint8), iterations=thickness) > 0) + + return edge_mask + + +def disk_kernel(radius: int) -> np.ndarray: + """ + Generate disk kernel with given radius. + + Args: + radius (int): Radius of the disk (in pixels). + + Returns: + np.ndarray: (2*radius+1, 2*radius+1) normalized convolution kernel. + """ + # Create coordinate grid centered at (0,0) + L = np.arange(-radius, radius + 1) + X, Y = np.meshgrid(L, L) + # Generate disk: region inside circle with radius R is 1 + kernel = ((X**2 + Y**2) <= radius**2).astype(np.float32) + # Normalize the kernel + kernel /= np.sum(kernel) + return kernel + + +def disk_blur(image: np.ndarray, radius: int) -> np.ndarray: + """ + Apply disk blur to an image using FFT convolution. + + Args: + image (np.ndarray): Input image, can be grayscale or color. + radius (int): Blur radius (in pixels). + + Returns: + np.ndarray: Blurred image. + """ + if radius == 0: + return image + kernel = disk_kernel(radius) + if image.ndim == 2: + blurred = fftconvolve(image, kernel, mode='same') + elif image.ndim == 3: + channels = [] + for i in range(image.shape[2]): + blurred_channel = fftconvolve(image[..., i], kernel, mode='same') + channels.append(blurred_channel) + blurred = np.stack(channels, axis=-1) + else: + raise ValueError("Image must be 2D or 3D.") + return blurred + + +def depth_of_field( + img: np.ndarray, + disp: np.ndarray, + focus_disp : float, + max_blur_radius : int = 10, +) -> np.ndarray: + """ + Apply depth of field effect to an image. + + Args: + img (numpy.ndarray): (H, W, 3) input image. + depth (numpy.ndarray): (H, W) depth map of the scene. + focus_depth (float): Focus depth of the lens. + strength (float): Strength of the depth of field effect. + max_blur_radius (int): Maximum blur radius (in pixels). + + Returns: + numpy.ndarray: (H, W, 3) output image with depth of field effect applied. + """ + # Precalculate dialated depth map for each blur radius + max_disp = np.max(disp) + disp = disp / max_disp + focus_disp = focus_disp / max_disp + dilated_disp = [] + for radius in range(max_blur_radius + 1): + dilated_disp.append(cv2.dilate(disp, cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2*radius+1, 2*radius+1)), iterations=1)) + + # Determine the blur radius for each pixel based on the depth map + blur_radii = np.clip(abs(disp - focus_disp) * max_blur_radius, 0, max_blur_radius).astype(np.int32) + for radius in range(max_blur_radius + 1): + dialted_blur_radii = np.clip(abs(dilated_disp[radius] - focus_disp) * max_blur_radius, 0, max_blur_radius).astype(np.int32) + mask = (dialted_blur_radii >= radius) & (dialted_blur_radii >= blur_radii) & (dilated_disp[radius] > disp) + blur_radii[mask] = dialted_blur_radii[mask] + blur_radii = np.clip(blur_radii, 0, max_blur_radius) + blur_radii = cv2.blur(blur_radii, (5, 5)) + + # Precalculate the blured image for each blur radius + unique_radii = np.unique(blur_radii) + precomputed = {} + for radius in range(max_blur_radius + 1): + if radius not in unique_radii: + continue + precomputed[radius] = disk_blur(img, radius) + + # Composit the blured image for each pixel + output = np.zeros_like(img) + for r in unique_radii: + mask = blur_radii == r + output[mask] = precomputed[r][mask] + + return output diff --git a/moge/utils/geometry_torch.py b/moge/utils/geometry_torch.py new file mode 100644 index 0000000..20b5632 --- /dev/null +++ b/moge/utils/geometry_torch.py @@ -0,0 +1,234 @@ +from typing import * +import math +from collections import namedtuple + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.types +import utils3d + +from .tools import timeit +from .geometry_numpy import solve_optimal_focal_shift, solve_optimal_shift + + +def weighted_mean(x: torch.Tensor, w: torch.Tensor = None, dim: Union[int, torch.Size] = None, keepdim: bool = False, eps: float = 1e-7) -> torch.Tensor: + if w is None: + return x.mean(dim=dim, keepdim=keepdim) + else: + w = w.to(x.dtype) + return (x * w).mean(dim=dim, keepdim=keepdim) / w.mean(dim=dim, keepdim=keepdim).add(eps) + + +def harmonic_mean(x: torch.Tensor, w: torch.Tensor = None, dim: Union[int, torch.Size] = None, keepdim: bool = False, eps: float = 1e-7) -> torch.Tensor: + if w is None: + return x.add(eps).reciprocal().mean(dim=dim, keepdim=keepdim).reciprocal() + else: + w = w.to(x.dtype) + return weighted_mean(x.add(eps).reciprocal(), w, dim=dim, keepdim=keepdim, eps=eps).add(eps).reciprocal() + + +def geometric_mean(x: torch.Tensor, w: torch.Tensor = None, dim: Union[int, torch.Size] = None, keepdim: bool = False, eps: float = 1e-7) -> torch.Tensor: + if w is None: + return x.add(eps).log().mean(dim=dim).exp() + else: + w = w.to(x.dtype) + return weighted_mean(x.add(eps).log(), w, dim=dim, keepdim=keepdim, eps=eps).exp() + + +def normalized_view_plane_uv(width: int, height: int, aspect_ratio: float = None, dtype: torch.dtype = None, device: torch.device = None) -> torch.Tensor: + "UV with left-top corner as (-width / diagonal, -height / diagonal) and right-bottom corner as (width / diagonal, height / diagonal)" + if aspect_ratio is None: + aspect_ratio = width / height + + span_x = aspect_ratio / (1 + aspect_ratio ** 2) ** 0.5 + span_y = 1 / (1 + aspect_ratio ** 2) ** 0.5 + + u = torch.linspace(-span_x * (width - 1) / width, span_x * (width - 1) / width, width, dtype=dtype, device=device) + v = torch.linspace(-span_y * (height - 1) / height, span_y * (height - 1) / height, height, dtype=dtype, device=device) + u, v = torch.meshgrid(u, v, indexing='xy') + uv = torch.stack([u, v], dim=-1) + return uv + + +def gaussian_blur_2d(input: torch.Tensor, kernel_size: int, sigma: float) -> torch.Tensor: + kernel = torch.exp(-(torch.arange(-kernel_size // 2 + 1, kernel_size // 2 + 1, dtype=input.dtype, device=input.device) ** 2) / (2 * sigma ** 2)) + kernel = kernel / kernel.sum() + kernel = (kernel[:, None] * kernel[None, :]).reshape(1, 1, kernel_size, kernel_size) + input = F.pad(input, (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2), mode='replicate') + input = F.conv2d(input, kernel, groups=input.shape[1]) + return input + + +def focal_to_fov(focal: torch.Tensor): + return 2 * torch.atan(0.5 / focal) + + +def fov_to_focal(fov: torch.Tensor): + return 0.5 / torch.tan(fov / 2) + + +def angle_diff_vec3(v1: torch.Tensor, v2: torch.Tensor, eps: float = 1e-12): + return torch.atan2(torch.cross(v1, v2, dim=-1).norm(dim=-1) + eps, (v1 * v2).sum(dim=-1)) + +def intrinsics_to_fov(intrinsics: torch.Tensor): + """ + Returns field of view in radians from normalized intrinsics matrix. + ### Parameters: + - intrinsics: torch.Tensor of shape (..., 3, 3) + + ### Returns: + - fov_x: torch.Tensor of shape (...) + - fov_y: torch.Tensor of shape (...) + """ + focal_x = intrinsics[..., 0, 0] + focal_y = intrinsics[..., 1, 1] + return 2 * torch.atan(0.5 / focal_x), 2 * torch.atan(0.5 / focal_y) + + +def point_map_to_depth_legacy(points: torch.Tensor): + height, width = points.shape[-3:-1] + diagonal = (height ** 2 + width ** 2) ** 0.5 + uv = normalized_view_plane_uv(width, height, dtype=points.dtype, device=points.device) # (H, W, 2) + + # Solve least squares problem + b = (uv * points[..., 2:]).flatten(-3, -1) # (..., H * W * 2) + A = torch.stack([points[..., :2], -uv.expand_as(points[..., :2])], dim=-1).flatten(-4, -2) # (..., H * W * 2, 2) + + M = A.transpose(-2, -1) @ A + solution = (torch.inverse(M + 1e-6 * torch.eye(2).to(A)) @ (A.transpose(-2, -1) @ b[..., None])).squeeze(-1) + focal, shift = solution.unbind(-1) + + depth = points[..., 2] + shift[..., None, None] + fov_x = torch.atan(width / diagonal / focal) * 2 + fov_y = torch.atan(height / diagonal / focal) * 2 + return depth, fov_x, fov_y, shift + + +def view_plane_uv_to_focal(uv: torch.Tensor): + normed_uv = normalized_view_plane_uv(width=uv.shape[-2], height=uv.shape[-3], device=uv.device, dtype=uv.dtype) + focal = (uv * normed_uv).sum() / uv.square().sum().add(1e-12) + return focal + + +def recover_focal_shift(points: torch.Tensor, mask: torch.Tensor = None, focal: torch.Tensor = None, downsample_size: Tuple[int, int] = (64, 64)): + """ + Recover the depth map and FoV from a point map with unknown z shift and focal. + + Note that it assumes: + - the optical center is at the center of the map + - the map is undistorted + - the map is isometric in the x and y directions + + ### Parameters: + - `points: torch.Tensor` of shape (..., H, W, 3) + - `downsample_size: Tuple[int, int]` in (height, width), the size of the downsampled map. Downsampling produces approximate solution and is efficient for large maps. + + ### Returns: + - `focal`: torch.Tensor of shape (...) the estimated focal length, relative to the half diagonal of the map + - `shift`: torch.Tensor of shape (...) Z-axis shift to translate the point map to camera space + """ + shape = points.shape + height, width = points.shape[-3], points.shape[-2] + diagonal = (height ** 2 + width ** 2) ** 0.5 + + points = points.reshape(-1, *shape[-3:]) + mask = None if mask is None else mask.reshape(-1, *shape[-3:-1]) + focal = focal.reshape(-1) if focal is not None else None + uv = normalized_view_plane_uv(width, height, dtype=points.dtype, device=points.device) # (H, W, 2) + + points_lr = F.interpolate(points.permute(0, 3, 1, 2), downsample_size, mode='nearest').permute(0, 2, 3, 1) + uv_lr = F.interpolate(uv.unsqueeze(0).permute(0, 3, 1, 2), downsample_size, mode='nearest').squeeze(0).permute(1, 2, 0) + mask_lr = None if mask is None else F.interpolate(mask.to(torch.float32).unsqueeze(1), downsample_size, mode='nearest').squeeze(1) > 0 + + uv_lr_np = uv_lr.cpu().numpy() + points_lr_np = points_lr.detach().cpu().numpy() + focal_np = focal.cpu().numpy() if focal is not None else None + mask_lr_np = None if mask is None else mask_lr.cpu().numpy() + optim_shift, optim_focal = [], [] + for i in range(points.shape[0]): + points_lr_i_np = points_lr_np[i] if mask is None else points_lr_np[i][mask_lr_np[i]] + uv_lr_i_np = uv_lr_np if mask is None else uv_lr_np[mask_lr_np[i]] + if uv_lr_i_np.shape[0] < 2: + optim_focal.append(1) + optim_shift.append(0) + continue + if focal is None: + optim_shift_i, optim_focal_i = solve_optimal_focal_shift(uv_lr_i_np, points_lr_i_np) + optim_focal.append(float(optim_focal_i)) + else: + optim_shift_i = solve_optimal_shift(uv_lr_i_np, points_lr_i_np, focal_np[i]) + optim_shift.append(float(optim_shift_i)) + optim_shift = torch.tensor(optim_shift, device=points.device, dtype=points.dtype).reshape(shape[:-3]) + + if focal is None: + optim_focal = torch.tensor(optim_focal, device=points.device, dtype=points.dtype).reshape(shape[:-3]) + else: + optim_focal = focal.reshape(shape[:-3]) + + return optim_focal, optim_shift + + +def theshold_depth_change(depth: torch.Tensor, mask: torch.Tensor, pooler: Literal['min', 'max'], rtol: float = 0.2, kernel_size: int = 3): + *batch_shape, height, width = depth.shape + depth = depth.reshape(-1, 1, height, width) + mask = mask.reshape(-1, 1, height, width) + if pooler =='max': + pooled_depth = F.max_pool2d(torch.where(mask, depth, -torch.inf), kernel_size, stride=1, padding=kernel_size // 2) + output_mask = pooled_depth > depth * (1 + rtol) + elif pooler =='min': + pooled_depth = -F.max_pool2d(-torch.where(mask, depth, torch.inf), kernel_size, stride=1, padding=kernel_size // 2) + output_mask = pooled_depth < depth * (1 - rtol) + else: + raise ValueError(f'Unsupported pooler: {pooler}') + output_mask = output_mask.reshape(*batch_shape, height, width) + return output_mask + + +def dilate_with_mask(input: torch.Tensor, mask: torch.BoolTensor, filter: Literal['min', 'max', 'mean', 'median'] = 'mean', iterations: int = 1) -> torch.Tensor: + kernel = torch.tensor([[False, True, False], [True, True, True], [False, True, False]], device=input.device, dtype=torch.bool) + for _ in range(iterations): + input_window = utils3d.pt.sliding_window(F.pad(input, (1, 1, 1, 1), mode='constant', value=0), window_size=3, stride=1, dim=(-2, -1)) + mask_window = kernel & utils3d.pt.sliding_window(F.pad(mask, (1, 1, 1, 1), mode='constant', value=False), window_size=3, stride=1, dim=(-2, -1)) + if filter =='min': + input = torch.where(mask, input, torch.where(mask_window, input_window, torch.inf).min(dim=(-2, -1)).values) + elif filter =='max': + input = torch.where(mask, input, torch.where(mask_window, input_window, -torch.inf).max(dim=(-2, -1)).values) + elif filter == 'mean': + input = torch.where(mask, input, torch.where(mask_window, input_window, torch.nan).nanmean(dim=(-2, -1))) + elif filter =='median': + input = torch.where(mask, input, torch.where(mask_window, input_window, torch.nan).flatten(-2).nanmedian(dim=-1).values) + mask = mask_window.any(dim=(-2, -1)) + return input, mask + + +def refine_depth_with_normal(depth: torch.Tensor, normal: torch.Tensor, intrinsics: torch.Tensor, iterations: int = 10, damp: float = 1e-3, eps: float = 1e-12, kernel_size: int = 5) -> torch.Tensor: + device, dtype = depth.device, depth.dtype + height, width = depth.shape[-2:] + radius = kernel_size // 2 + + duv = torch.stack(torch.meshgrid(torch.linspace(-radius / width, radius / width, kernel_size, device=device, dtype=dtype), torch.linspace(-radius / height, radius / height, kernel_size, device=device, dtype=dtype), indexing='xy'), dim=-1).to(dtype=dtype, device=device) + + log_depth = depth.clamp_min_(eps).log() + log_depth_diff = utils3d.pt.sliding_window(log_depth, window_size=kernel_size, stride=1, dim=(-2, -1)) - log_depth[..., radius:-radius, radius:-radius, None, None] + + weight = torch.exp(-(log_depth_diff / duv.norm(dim=-1).clamp_min_(eps) / 10).square()) + tot_weight = weight.sum(dim=(-2, -1)).clamp_min_(eps) + + uv = utils3d.pt.uv_map((height, width), device=device, dtype=dtype) + K_inv = torch.inverse(intrinsics) + + grad = -(normal[..., None, :2] @ K_inv[..., None, None, :2, :2]).squeeze(-2) \ + / (normal[..., None, 2:] + normal[..., None, :2] @ (K_inv[..., None, None, :2, :2] @ uv[..., :, None] + K_inv[..., None, None, :2, 2:])).squeeze(-2) + laplacian = (weight * ((utils3d.pt.sliding_window(grad, window_size=kernel_size, stride=1, dim=(-3, -2)) + grad[..., radius:-radius, radius:-radius, :, None, None]) * (duv.permute(2, 0, 1) / 2)).sum(dim=-3)).sum(dim=(-2, -1)) + + laplacian = laplacian.clamp(-0.1, 0.1) + log_depth_refine = log_depth.clone() + + for _ in range(iterations): + log_depth_refine[..., radius:-radius, radius:-radius] = 0.1 * log_depth_refine[..., radius:-radius, radius:-radius] + 0.9 * (damp * log_depth[..., radius:-radius, radius:-radius] - laplacian + (weight * utils3d.pt.sliding_window_2d(log_depth_refine, window_size=kernel_size, stride=1, dim=(-2, -1))).sum(dim=(-2, -1))) / (tot_weight + damp) + + depth_refine = log_depth_refine.exp() + + return depth_refine \ No newline at end of file diff --git a/moge/utils/io.py b/moge/utils/io.py new file mode 100644 index 0000000..47b1641 --- /dev/null +++ b/moge/utils/io.py @@ -0,0 +1,271 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +from typing import IO +import zipfile +import json +import io +from typing import * +from pathlib import Path +import re +from PIL import Image, PngImagePlugin + +import numpy as np +import cv2 + +from .tools import timeit + + +def save_glb( + save_path: Union[str, os.PathLike], + vertices: np.ndarray, + faces: np.ndarray, + vertex_uvs: np.ndarray, + texture: np.ndarray, + vertex_normals: Optional[np.ndarray] = None, +): + import trimesh + import trimesh.visual + from PIL import Image + + trimesh.Trimesh( + vertices=vertices, + vertex_normals=vertex_normals, + faces=faces, + visual = trimesh.visual.texture.TextureVisuals( + uv=vertex_uvs, + material=trimesh.visual.material.PBRMaterial( + baseColorTexture=Image.fromarray(texture), + metallicFactor=0.5, + roughnessFactor=1.0 + ) + ), + process=False + ).export(save_path) + + +def save_ply( + save_path: Union[str, os.PathLike], + vertices: np.ndarray, + faces: np.ndarray, + vertex_colors: np.ndarray, + vertex_normals: Optional[np.ndarray] = None, +): + import trimesh + import trimesh.visual + from PIL import Image + + trimesh.Trimesh( + vertices=vertices, + faces=faces, + vertex_colors=vertex_colors, + vertex_normals=vertex_normals, + process=False + ).export(save_path) + + +def read_image(path: Union[str, os.PathLike, IO]) -> np.ndarray: + """ + Read a image, return uint8 RGB array of shape (H, W, 3). + """ + if isinstance(path, (str, os.PathLike)): + data = Path(path).read_bytes() + else: + data = path.read() + image = cv2.cvtColor(cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_COLOR), cv2.COLOR_BGR2RGB) + return image + + +def write_image(path: Union[str, os.PathLike, IO], image: np.ndarray, quality: int = 95): + """ + Write a image, input uint8 RGB array of shape (H, W, 3). + """ + data = cv2.imencode('.jpg', cv2.cvtColor(image, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_JPEG_QUALITY, quality])[1].tobytes() + if isinstance(path, (str, os.PathLike)): + Path(path).write_bytes(data) + else: + path.write(data) + + +def read_depth(path: Union[str, os.PathLike, IO]) -> np.ndarray: + """ + Read a depth image, return float32 depth array of shape (H, W). + """ + if isinstance(path, (str, os.PathLike)): + data = Path(path).read_bytes() + else: + data = path.read() + pil_image = Image.open(io.BytesIO(data)) + near = float(pil_image.info.get('near')) + far = float(pil_image.info.get('far')) + depth = np.array(pil_image) + mask_nan, mask_inf = depth == 0, depth == 65535 + depth = (depth.astype(np.float32) - 1) / 65533 + depth = near ** (1 - depth) * far ** depth + if 'unit' in pil_image.info: # Legacy support for depth units + unit = float(pil_image.info.get('unit')) + depth = depth * unit + depth[mask_nan] = np.nan + depth[mask_inf] = np.inf + return depth + + +def write_depth( + path: Union[str, os.PathLike, IO], + depth: np.ndarray, + max_range: float = 1e5, + compression_level: int = 7, +): + """ + Encode and write a depth image as 16-bit PNG format. + ## Parameters: + - `path: Union[str, os.PathLike, IO]` + The file path or file object to write to. + - `depth: np.ndarray` + The depth array, float32 array of shape (H, W). + May contain `NaN` for invalid values and `Inf` for infinite values. + + Depth values are encoded as follows: + - 0: unknown + - 1 ~ 65534: depth values in logarithmic + - 65535: infinity + + metadata is stored in the PNG file as text fields: + - `near`: the minimum depth value + - `far`: the maximum depth value + """ + mask_values, mask_nan, mask_inf = np.isfinite(depth), np.isnan(depth),np.isinf(depth) + + depth = depth.astype(np.float32) + mask_finite = depth + near = max(depth[mask_values].min(), 1e-5) + far = max(near * 1.1, min(depth[mask_values].max(), near * max_range)) + depth = 1 + np.round((np.log(np.nan_to_num(depth, nan=0).clip(near, far) / near) / np.log(far / near)).clip(0, 1) * 65533).astype(np.uint16) # 1~65534 + depth[mask_nan] = 0 + depth[mask_inf] = 65535 + + pil_image = Image.fromarray(depth) + pnginfo = PngImagePlugin.PngInfo() + pnginfo.add_text('near', str(near)) + pnginfo.add_text('far', str(far)) + pil_image.save(path, pnginfo=pnginfo, compress_level=compression_level) + + +def read_segmentation(path: Union[str, os.PathLike, IO]) -> Tuple[np.ndarray, Dict[str, int]]: + """ + Read a segmentation mask + ### Parameters: + - `path: Union[str, os.PathLike, IO]` + The file path or file object to read from. + ### Returns: + - `Tuple[np.ndarray, Dict[str, int]]` + A tuple containing: + - `mask`: uint8 or uint16 numpy.ndarray of shape (H, W). + - `labels`: Dict[str, int]. The label mapping, a dictionary of {label_name: label_id}. + """ + if isinstance(path, (str, os.PathLike)): + data = Path(path).read_bytes() + else: + data = path.read() + pil_image = Image.open(io.BytesIO(data)) + labels = json.loads(pil_image.info['labels']) if 'labels' in pil_image.info else None + mask = np.array(pil_image) + return mask, labels + + +def write_segmentation(path: Union[str, os.PathLike, IO], mask: np.ndarray, labels: Dict[str, int] = None, compression_level: int = 7): + """ + Write a segmentation mask and label mapping, as PNG format. + ### Parameters: + - `path: Union[str, os.PathLike, IO]` + The file path or file object to write to. + - `mask: np.ndarray` + The segmentation mask, uint8 or uint16 array of shape (H, W). + - `labels: Dict[str, int] = None` + The label mapping, a dictionary of {label_name: label_id}. + - `compression_level: int = 7` + The compression level for PNG compression. + """ + assert mask.dtype == np.uint8 or mask.dtype == np.uint16, f"Unsupported dtype {mask.dtype}" + pil_image = Image.fromarray(mask) + pnginfo = PngImagePlugin.PngInfo() + if labels is not None: + labels_json = json.dumps(labels, ensure_ascii=True, separators=(',', ':')) + pnginfo.add_text('labels', labels_json) + pil_image.save(path, pnginfo=pnginfo, compress_level=compression_level) + + + +def read_normal(path: Union[str, os.PathLike, IO]) -> np.ndarray: + """ + Read a normal image, return float32 normal array of shape (H, W, 3). + """ + if isinstance(path, (str, os.PathLike)): + data = Path(path).read_bytes() + else: + data = path.read() + normal = cv2.cvtColor(cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_UNCHANGED), cv2.COLOR_BGR2RGB) + mask_nan = np.all(normal == 0, axis=-1) + normal = (normal.astype(np.float32) / 65535 - 0.5) * [2.0, -2.0, -2.0] + normal = normal / (np.sqrt(np.square(normal[..., 0]) + np.square(normal[..., 1]) + np.square(normal[..., 2])) + 1e-12) + normal[mask_nan] = np.nan + return normal + + +def write_normal(path: Union[str, os.PathLike, IO], normal: np.ndarray, compression_level: int = 7) -> np.ndarray: + """ + Write a normal image, input float32 normal array of shape (H, W, 3). + """ + mask_nan = np.isnan(normal).any(axis=-1) + normal = ((normal * [0.5, -0.5, -0.5] + 0.5).clip(0, 1) * 65535).astype(np.uint16) + normal[mask_nan] = 0 + data = cv2.imencode('.png', cv2.cvtColor(normal, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_PNG_COMPRESSION, compression_level])[1].tobytes() + if isinstance(path, (str, os.PathLike)): + Path(path).write_bytes(data) + else: + path.write(data) + + +def read_mask(path: Union[str, os.PathLike, IO[bytes]]) -> np.ndarray: + """ + Read a binary mask, return bool array of shape (H, W). + """ + if isinstance(path, (str, os.PathLike)): + data = Path(path).read_bytes() + else: + data = path.read() + mask = cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_UNCHANGED) + if len(mask.shape) == 3: + mask = mask[..., 0] + return mask > 0 + + +def write_mask(path: Union[str, os.PathLike, IO[bytes]], mask: np.ndarray, compression_level: int = 7): + """ + Write a binary mask, input bool array of shape (H, W). + """ + assert mask.dtype == bool, f"Mask must be bool array, got {mask.dtype}" + mask = (mask.astype(np.uint8) * 255).astype(np.uint8) + data = cv2.imencode('.png', mask, [cv2.IMWRITE_PNG_COMPRESSION, compression_level])[1].tobytes() + if isinstance(path, (str, os.PathLike)): + Path(path).write_bytes(data) + else: + path.write(data) + + +JSON_TYPE = Union[str, int, float, bool, None, Dict[str, "JSON"], List["JSON"]] + + +def read_json(path: Union[str, os.PathLike, IO[str]]) -> JSON_TYPE: + if isinstance(path, (str, os.PathLike)): + text = Path(path).read_text() + else: + text = path.read() + return json.loads(text) + + +def write_json(path: Union[str, os.PathLike, IO[str]], content: JSON_TYPE): + text = json.dumps(content) + if isinstance(path, (str, os.PathLike)): + Path(path).write_text(text) + else: + path.write(text) \ No newline at end of file diff --git a/moge/utils/panorama.py b/moge/utils/panorama.py new file mode 100644 index 0000000..42d915a --- /dev/null +++ b/moge/utils/panorama.py @@ -0,0 +1,191 @@ +import os +os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1' +from pathlib import Path +from typing import * +import itertools +import json +import warnings + +import cv2 +import numpy as np +from numpy import ndarray +from tqdm import tqdm, trange +from scipy.sparse import csr_array, hstack, vstack +from scipy.ndimage import convolve +from scipy.sparse.linalg import lsmr + +import utils3d + + +def get_panorama_cameras(): + vertices, _ = utils3d.np.create_icosahedron_mesh() + intrinsics = utils3d.np.intrinsics_from_fov(fov_x=np.deg2rad(90), fov_y=np.deg2rad(90)) + extrinsics = utils3d.np.extrinsics_look_at([0, 0, 0], vertices, [0, 0, 1]).astype(np.float32) + return extrinsics, [intrinsics] * len(vertices) + + +def spherical_uv_to_directions(uv: np.ndarray): + theta, phi = (1 - uv[..., 0]) * (2 * np.pi), uv[..., 1] * np.pi + directions = np.stack([np.sin(phi) * np.cos(theta), np.sin(phi) * np.sin(theta), np.cos(phi)], axis=-1) + return directions + + +def directions_to_spherical_uv(directions: np.ndarray): + directions = directions / np.linalg.norm(directions, axis=-1, keepdims=True) + u = 1 - np.arctan2(directions[..., 1], directions[..., 0]) / (2 * np.pi) % 1.0 + v = np.arccos(directions[..., 2]) / np.pi + return np.stack([u, v], axis=-1) + + +def split_panorama_image(image: np.ndarray, extrinsics: np.ndarray, intrinsics: np.ndarray, resolution: int): + height, width = image.shape[:2] + uv = utils3d.np.uv_map((resolution, resolution)) + splitted_images = [] + for i in range(len(extrinsics)): + spherical_uv = directions_to_spherical_uv(utils3d.np.unproject_cv(uv, np.ones_like(uv[..., 0]), extrinsics=extrinsics[i], intrinsics=intrinsics[i])) + pixels = utils3d.np.uv_to_pixel(spherical_uv, (height, width)).astype(np.float32) + + splitted_image = cv2.remap(image, pixels[..., 0], pixels[..., 1], interpolation=cv2.INTER_LINEAR) + splitted_images.append(splitted_image) + return splitted_images + + +def poisson_equation(width: int, height: int, wrap_x: bool = False, wrap_y: bool = False) -> Tuple[csr_array, ndarray]: + grid_index = np.arange(height * width).reshape(height, width) + grid_index = np.pad(grid_index, ((0, 0), (1, 1)), mode='wrap' if wrap_x else 'edge') + grid_index = np.pad(grid_index, ((1, 1), (0, 0)), mode='wrap' if wrap_y else 'edge') + + data = np.array([[-4, 1, 1, 1, 1]], dtype=np.float32).repeat(height * width, axis=0).reshape(-1) + indices = np.stack([ + grid_index[1:-1, 1:-1], + grid_index[:-2, 1:-1], # up + grid_index[2:, 1:-1], # down + grid_index[1:-1, :-2], # left + grid_index[1:-1, 2:] # right + ], axis=-1).reshape(-1) + indptr = np.arange(0, height * width * 5 + 1, 5) + A = csr_array((data, indices, indptr), shape=(height * width, height * width)) + + return A + + +def grad_equation(width: int, height: int, wrap_x: bool = False, wrap_y: bool = False) -> Tuple[csr_array, np.ndarray]: + grid_index = np.arange(width * height).reshape(height, width) + if wrap_x: + grid_index = np.pad(grid_index, ((0, 0), (0, 1)), mode='wrap') + if wrap_y: + grid_index = np.pad(grid_index, ((0, 1), (0, 0)), mode='wrap') + + data = np.concatenate([ + np.concatenate([ + np.ones((grid_index.shape[0], grid_index.shape[1] - 1), dtype=np.float32).reshape(-1, 1), # x[i,j] + -np.ones((grid_index.shape[0], grid_index.shape[1] - 1), dtype=np.float32).reshape(-1, 1), # x[i,j-1] + ], axis=1).reshape(-1), + np.concatenate([ + np.ones((grid_index.shape[0] - 1, grid_index.shape[1]), dtype=np.float32).reshape(-1, 1), # x[i,j] + -np.ones((grid_index.shape[0] - 1, grid_index.shape[1]), dtype=np.float32).reshape(-1, 1), # x[i-1,j] + ], axis=1).reshape(-1), + ]) + indices = np.concatenate([ + np.concatenate([ + grid_index[:, :-1].reshape(-1, 1), + grid_index[:, 1:].reshape(-1, 1), + ], axis=1).reshape(-1), + np.concatenate([ + grid_index[:-1, :].reshape(-1, 1), + grid_index[1:, :].reshape(-1, 1), + ], axis=1).reshape(-1), + ]) + indptr = np.arange(0, grid_index.shape[0] * (grid_index.shape[1] - 1) * 2 + (grid_index.shape[0] - 1) * grid_index.shape[1] * 2 + 1, 2) + A = csr_array((data, indices, indptr), shape=(grid_index.shape[0] * (grid_index.shape[1] - 1) + (grid_index.shape[0] - 1) * grid_index.shape[1], height * width)) + + return A + + +def merge_panorama_depth(width: int, height: int, distance_maps: List[np.ndarray], pred_masks: List[np.ndarray], extrinsics: List[np.ndarray], intrinsics: List[np.ndarray]): + if max(width, height) > 256: + panorama_depth_init, _ = merge_panorama_depth(width // 2, height // 2, distance_maps, pred_masks, extrinsics, intrinsics) + panorama_depth_init = cv2.resize(panorama_depth_init, (width, height), cv2.INTER_LINEAR) + else: + panorama_depth_init = None + + uv = utils3d.np.uv_map(height, width) + spherical_directions = spherical_uv_to_directions(uv) + + # Warp each view to the panorama + panorama_log_distance_grad_maps, panorama_grad_masks = [], [] + panorama_log_distance_laplacian_maps, panorama_laplacian_masks = [], [] + panorama_pred_masks = [] + for i in range(len(distance_maps)): + projected_uv, projected_depth = utils3d.np.project_cv(spherical_directions, extrinsics=extrinsics[i], intrinsics=intrinsics[i]) + projection_valid_mask = (projected_depth > 0) & (projected_uv > 0).all(axis=-1) & (projected_uv < 1).all(axis=-1) + + projected_pixels = utils3d.np.uv_to_pixel(np.clip(projected_uv, 0, 1), distance_maps[i].shape).astype(np.float32) + + log_splitted_distance = np.log(distance_maps[i]) + panorama_log_distance_map = np.where(projection_valid_mask, cv2.remap(log_splitted_distance, projected_pixels[..., 0], projected_pixels[..., 1], cv2.INTER_LINEAR, borderMode=cv2.BORDER_REPLICATE), 0) + panorama_pred_mask = projection_valid_mask & (cv2.remap(pred_masks[i].astype(np.uint8), projected_pixels[..., 0], projected_pixels[..., 1], cv2.INTER_NEAREST, borderMode=cv2.BORDER_REPLICATE) > 0) + + # calculate gradient map + padded = np.pad(panorama_log_distance_map, ((0, 0), (0, 1)), mode='wrap') + grad_x, grad_y = padded[:, :-1] - padded[:, 1:], padded[:-1, :] - padded[1:, :] + + padded = np.pad(panorama_pred_mask, ((0, 0), (0, 1)), mode='wrap') + mask_x, mask_y = padded[:, :-1] & padded[:, 1:], padded[:-1, :] & padded[1:, :] + + panorama_log_distance_grad_maps.append((grad_x, grad_y)) + panorama_grad_masks.append((mask_x, mask_y)) + + # calculate laplacian map + padded = np.pad(panorama_log_distance_map, ((1, 1), (0, 0)), mode='edge') + padded = np.pad(padded, ((0, 0), (1, 1)), mode='wrap') + laplacian = convolve(padded, np.array([[0, 1, 0], [1, -4, 1], [0, 1, 0]], dtype=np.float32))[1:-1, 1:-1] + + padded = np.pad(panorama_pred_mask, ((1, 1), (0, 0)), mode='edge') + padded = np.pad(padded, ((0, 0), (1, 1)), mode='wrap') + mask = convolve(padded.astype(np.uint8), np.array([[0, 1, 0], [1, 1, 1], [0, 1, 0]], dtype=np.uint8))[1:-1, 1:-1] == 5 + + panorama_log_distance_laplacian_maps.append(laplacian) + panorama_laplacian_masks.append(mask) + + panorama_pred_masks.append(panorama_pred_mask) + + panorama_log_distance_grad_x = np.stack([grad_map[0] for grad_map in panorama_log_distance_grad_maps], axis=0) + panorama_log_distance_grad_y = np.stack([grad_map[1] for grad_map in panorama_log_distance_grad_maps], axis=0) + panorama_grad_mask_x = np.stack([mask_map[0] for mask_map in panorama_grad_masks], axis=0) + panorama_grad_mask_y = np.stack([mask_map[1] for mask_map in panorama_grad_masks], axis=0) + + panorama_log_distance_grad_x = np.sum(panorama_log_distance_grad_x * panorama_grad_mask_x, axis=0) / np.sum(panorama_grad_mask_x, axis=0).clip(1e-3) + panorama_log_distance_grad_y = np.sum(panorama_log_distance_grad_y * panorama_grad_mask_y, axis=0) / np.sum(panorama_grad_mask_y, axis=0).clip(1e-3) + + panorama_laplacian_maps = np.stack(panorama_log_distance_laplacian_maps, axis=0) + panorama_laplacian_masks = np.stack(panorama_laplacian_masks, axis=0) + panorama_laplacian_map = np.sum(panorama_laplacian_maps * panorama_laplacian_masks, axis=0) / np.sum(panorama_laplacian_masks, axis=0).clip(1e-3) + + grad_x_mask = np.any(panorama_grad_mask_x, axis=0).reshape(-1) + grad_y_mask = np.any(panorama_grad_mask_y, axis=0).reshape(-1) + grad_mask = np.concatenate([grad_x_mask, grad_y_mask]) + laplacian_mask = np.any(panorama_laplacian_masks, axis=0).reshape(-1) + + # Solve overdetermined system + A = vstack([ + grad_equation(width, height, wrap_x=True, wrap_y=False)[grad_mask], + poisson_equation(width, height, wrap_x=True, wrap_y=False)[laplacian_mask], + ]) + b = np.concatenate([ + panorama_log_distance_grad_x.reshape(-1)[grad_x_mask], + panorama_log_distance_grad_y.reshape(-1)[grad_y_mask], + panorama_laplacian_map.reshape(-1)[laplacian_mask] + ]) + x, *_ = lsmr( + A, b, + atol=1e-5, btol=1e-5, + x0=np.log(panorama_depth_init).reshape(-1) if panorama_depth_init is not None else None, + show=False, + ) + + panorama_depth = np.exp(x).reshape(height, width).astype(np.float32) + panorama_mask = np.any(panorama_pred_masks, axis=0) + + return panorama_depth, panorama_mask + diff --git a/moge/utils/tools.py b/moge/utils/tools.py new file mode 100644 index 0000000..3687f69 --- /dev/null +++ b/moge/utils/tools.py @@ -0,0 +1,289 @@ +from typing import * +import time +from pathlib import Path +from numbers import Number +from functools import wraps +import warnings +import math +import json +import os +import importlib +import importlib.util + + +def catch_exception(fn): + @wraps(fn) + def wrapper(*args, **kwargs): + try: + return fn(*args, **kwargs) + except Exception as e: + import traceback + print(f"Exception in {fn.__name__}", end='r') + # print({', '.join(repr(arg) for arg in args)}, {', '.join(f'{k}={v!r}' for k, v in kwargs.items())}) + traceback.print_exc(chain=False) + time.sleep(0.1) + return None + return wrapper + + +class CallbackOnException: + def __init__(self, callback: Callable, exception: type): + self.exception = exception + self.callback = callback + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + if isinstance(exc_val, self.exception): + self.callback() + return True + return False + +def traverse_nested_dict_keys(d: Dict[str, Dict]) -> Generator[Tuple[str, ...], None, None]: + for k, v in d.items(): + if isinstance(v, dict): + for sub_key in traverse_nested_dict_keys(v): + yield (k, ) + sub_key + else: + yield (k, ) + + +def get_nested_dict(d: Dict[str, Dict], keys: Tuple[str, ...], default: Any = None): + for k in keys: + d = d.get(k, default) + if d is None: + break + return d + +def set_nested_dict(d: Dict[str, Dict], keys: Tuple[str, ...], value: Any): + for k in keys[:-1]: + d = d.setdefault(k, {}) + d[keys[-1]] = value + + +def key_average(list_of_dicts: list) -> Dict[str, Any]: + """ + Returns a dictionary with the average value of each key in the input list of dictionaries. + """ + _nested_dict_keys = set() + for d in list_of_dicts: + _nested_dict_keys.update(traverse_nested_dict_keys(d)) + _nested_dict_keys = sorted(_nested_dict_keys) + result = {} + for k in _nested_dict_keys: + values = [] + for d in list_of_dicts: + v = get_nested_dict(d, k) + if v is not None and not math.isnan(v): + values.append(v) + avg = sum(values) / len(values) if values else float('nan') + set_nested_dict(result, k, avg) + return result + + +def flatten_nested_dict(d: Dict[str, Any], parent_key: Tuple[str, ...] = None) -> Dict[Tuple[str, ...], Any]: + """ + Flattens a nested dictionary into a single-level dictionary, with keys as tuples. + """ + items = [] + if parent_key is None: + parent_key = () + for k, v in d.items(): + new_key = parent_key + (k, ) + if isinstance(v, MutableMapping): + items.extend(flatten_nested_dict(v, new_key).items()) + else: + items.append((new_key, v)) + return dict(items) + + +def unflatten_nested_dict(d: Dict[str, Any]) -> Dict[str, Any]: + """ + Unflattens a single-level dictionary into a nested dictionary, with keys as tuples. + """ + result = {} + for k, v in d.items(): + sub_dict = result + for k_ in k[:-1]: + if k_ not in sub_dict: + sub_dict[k_] = {} + sub_dict = sub_dict[k_] + sub_dict[k[-1]] = v + return result + + +def read_jsonl(file): + import json + with open(file, 'r') as f: + data = f.readlines() + return [json.loads(line) for line in data] + + +def write_jsonl(data: List[dict], file): + import json + with open(file, 'w') as f: + for item in data: + f.write(json.dumps(item) + '\n') + + +def to_hierachical_dataframe(data: List[Dict[Tuple[str, ...], Any]]): + import pandas as pd + data = [flatten_nested_dict(d) for d in data] + df = pd.DataFrame(data) + df = df.sort_index(axis=1) + df.columns = pd.MultiIndex.from_tuples(df.columns) + return df + + +def recursive_replace(d: Union[List, Dict, str], mapping: Dict[str, str]): + if isinstance(d, str): + for old, new in mapping.items(): + d = d.replace(old, new) + elif isinstance(d, list): + for i, item in enumerate(d): + d[i] = recursive_replace(item, mapping) + elif isinstance(d, dict): + for k, v in d.items(): + d[k] = recursive_replace(v, mapping) + return d + + +class timeit: + _history: Dict[str, List['timeit']] = {} + + def __init__(self, name: str = None, verbose: bool = True, average: bool = False): + self.name = name + self.verbose = verbose + self.start = None + self.end = None + self.average = average + if average and name not in timeit._history: + timeit._history[name] = [] + + def __call__(self, func: Callable): + import inspect + if inspect.iscoroutinefunction(func): + async def wrapper(*args, **kwargs): + with timeit(self.name or func.__qualname__): + ret = await func(*args, **kwargs) + return ret + return wrapper + else: + def wrapper(*args, **kwargs): + with timeit(self.name or func.__qualname__): + ret = func(*args, **kwargs) + return ret + return wrapper + + def __enter__(self): + self.start = time.time() + return self + + @property + def time(self) -> float: + assert self.start is not None, "Time not yet started." + assert self.end is not None, "Time not yet ended." + return self.end - self.start + + @property + def average_time(self) -> float: + assert self.average, "Average time not available." + return sum(t.time for t in timeit._history[self.name]) / len(timeit._history[self.name]) + + @property + def history(self) -> List['timeit']: + return timeit._history.get(self.name, []) + + def __exit__(self, exc_type, exc_val, exc_tb): + self.end = time.time() + if self.average: + timeit._history[self.name].append(self) + if self.verbose: + if self.average: + avg = self.average_time + print(f"{self.name or 'It'} took {avg:.6f} seconds in average.") + else: + print(f"{self.name or 'It'} took {self.time:.6f} seconds.") + + +def strip_common_prefix_suffix(strings: List[str]) -> List[str]: + first = strings[0] + + for start in range(len(first)): + if any(s[start] != strings[0][start] for s in strings): + break + + for end in range(1, min(len(s) for s in strings)): + if any(s[-end] != first[-end] for s in strings): + break + + return [s[start:len(s) - end + 1] for s in strings] + + +def multithead_execute(inputs: List[Any], num_workers: int, pbar = None): + from concurrent.futures import ThreadPoolExecutor + from contextlib import nullcontext + from tqdm import tqdm + + if pbar is not None: + pbar.total = len(inputs) if hasattr(inputs, '__len__') else None + else: + pbar = tqdm(total=len(inputs) if hasattr(inputs, '__len__') else None) + + def decorator(fn: Callable): + with ( + ThreadPoolExecutor(max_workers=num_workers) as executor, + pbar + ): + pbar.refresh() + @catch_exception + @suppress_traceback + def _fn(input): + ret = fn(input) + pbar.update() + return ret + executor.map(_fn, inputs) + executor.shutdown(wait=True) + + return decorator + + +def suppress_traceback(fn): + @wraps(fn) + def wrapper(*args, **kwargs): + try: + return fn(*args, **kwargs) + except Exception as e: + e.__traceback__ = e.__traceback__.tb_next.tb_next + raise + return wrapper + + +class no_warnings: + def __init__(self, action: str = 'ignore', **kwargs): + self.action = action + self.filter_kwargs = kwargs + + def __call__(self, fn): + @wraps(fn) + def wrapper(*args, **kwargs): + with warnings.catch_warnings(): + warnings.simplefilter(self.action, **self.filter_kwargs) + return fn(*args, **kwargs) + return wrapper + + def __enter__(self): + self.warnings_manager = warnings.catch_warnings() + self.warnings_manager.__enter__() + warnings.simplefilter(self.action, **self.filter_kwargs) + + def __exit__(self, exc_type, exc_val, exc_tb): + self.warnings_manager.__exit__(exc_type, exc_val, exc_tb) + + +def import_file_as_module(file_path: Union[str, os.PathLike], module_name: str): + spec = importlib.util.spec_from_file_location(module_name, file_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module \ No newline at end of file diff --git a/moge/utils/vis.py b/moge/utils/vis.py new file mode 100644 index 0000000..cb9c237 --- /dev/null +++ b/moge/utils/vis.py @@ -0,0 +1,65 @@ +from typing import * + +import numpy as np +import matplotlib + + +def colorize_depth(depth: np.ndarray, mask: np.ndarray = None, normalize: bool = True, cmap: str = 'Spectral') -> np.ndarray: + if mask is None: + depth = np.where(depth > 0, depth, np.nan) + else: + depth = np.where((depth > 0) & mask, depth, np.nan) + disp = 1 / depth + if normalize: + min_disp, max_disp = np.nanquantile(disp, 0.001), np.nanquantile(disp, 0.99) + disp = (disp - min_disp) / (max_disp - min_disp) + colored = np.nan_to_num(matplotlib.colormaps[cmap](1.0 - disp)[..., :3], 0) + colored = np.ascontiguousarray((colored.clip(0, 1) * 255).astype(np.uint8)) + return colored + + +def colorize_depth_affine(depth: np.ndarray, mask: np.ndarray = None, cmap: str = 'Spectral') -> np.ndarray: + if mask is not None: + depth = np.where(mask, depth, np.nan) + + min_depth, max_depth = np.nanquantile(depth, 0.001), np.nanquantile(depth, 0.999) + depth = (depth - min_depth) / (max_depth - min_depth) + colored = np.nan_to_num(matplotlib.colormaps[cmap](depth)[..., :3], 0) + colored = np.ascontiguousarray((colored.clip(0, 1) * 255).astype(np.uint8)) + return colored + + +def colorize_disparity(disparity: np.ndarray, mask: np.ndarray = None, normalize: bool = True, cmap: str = 'Spectral') -> np.ndarray: + if mask is not None: + disparity = np.where(mask, disparity, np.nan) + + if normalize: + min_disp, max_disp = np.nanquantile(disparity, 0.001), np.nanquantile(disparity, 0.999) + disparity = (disparity - min_disp) / (max_disp - min_disp) + colored = np.nan_to_num(matplotlib.colormaps[cmap](1.0 - disparity)[..., :3], 0) + colored = np.ascontiguousarray((colored.clip(0, 1) * 255).astype(np.uint8)) + return colored + + +def colorize_segmentation(segmentation: np.ndarray, cmap: str = 'Set1') -> np.ndarray: + colored = matplotlib.colormaps[cmap]((segmentation % 20) / 20)[..., :3] + colored = np.ascontiguousarray((colored.clip(0, 1) * 255).astype(np.uint8)) + return colored + + +def colorize_normal(normal: np.ndarray, mask: np.ndarray = None) -> np.ndarray: + if mask is not None: + normal = np.where(mask[..., None], normal, 0) + normal = normal * [0.5, -0.5, -0.5] + 0.5 + normal = (normal.clip(0, 1) * 255).astype(np.uint8) + return normal + + +def colorize_error_map(error_map: np.ndarray, mask: np.ndarray = None, cmap: str = 'plasma', value_range: Tuple[float, float] = None): + vmin, vmax = value_range if value_range is not None else (np.nanmin(error_map), np.nanmax(error_map)) + cmap = matplotlib.colormaps[cmap] + colorized_error_map = cmap(((error_map - vmin) / (vmax - vmin)).clip(0, 1))[..., :3] + if mask is not None: + colorized_error_map = np.where(mask[..., None], colorized_error_map, 0) + colorized_error_map = np.ascontiguousarray((colorized_error_map.clip(0, 1) * 255).astype(np.uint8)) + return colorized_error_map diff --git a/moge/utils/webfile.py b/moge/utils/webfile.py new file mode 100644 index 0000000..1e98abf --- /dev/null +++ b/moge/utils/webfile.py @@ -0,0 +1,73 @@ +import requests +from typing import * + +__all__ = ["WebFile"] + + +class WebFile: + def __init__(self, url: str, session: Optional[requests.Session] = None, headers: Optional[Dict[str, str]] = None, size: Optional[int] = None): + self.url = url + self.session = session or requests.Session() + self.session.headers.update(headers or {}) + self._offset = 0 + self.size = size if size is not None else self._fetch_size() + + def _fetch_size(self): + with self.session.get(self.url, stream=True) as response: + response.raise_for_status() + content_length = response.headers.get("Content-Length") + if content_length is None: + raise ValueError("Missing Content-Length in header") + return int(content_length) + + def _fetch_data(self, offset: int, n: int) -> bytes: + headers = {"Range": f"bytes={offset}-{min(offset + n - 1, self.size)}"} + response = self.session.get(self.url, headers=headers) + response.raise_for_status() + return response.content + + def seekable(self) -> bool: + return True + + def tell(self) -> int: + return self._offset + + def available(self) -> int: + return self.size - self._offset + + def seek(self, offset: int, whence: int = 0) -> None: + if whence == 0: + new_offset = offset + elif whence == 1: + new_offset = self._offset + offset + elif whence == 2: + new_offset = self.size + offset + else: + raise ValueError("Invalid value for whence") + + self._offset = max(0, min(new_offset, self.size)) + + def read(self, n: Optional[int] = None) -> bytes: + if n is None or n < 0: + n = self.available() + else: + n = min(n, self.available()) + + if n == 0: + return b'' + + data = self._fetch_data(self._offset, n) + self._offset += len(data) + + return data + + def close(self) -> None: + pass + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + pass + + \ No newline at end of file diff --git a/moge/utils/webzipfile.py b/moge/utils/webzipfile.py new file mode 100644 index 0000000..25ed1d3 --- /dev/null +++ b/moge/utils/webzipfile.py @@ -0,0 +1,128 @@ +from typing import * +import io +import os +from zipfile import ( + ZipInfo, BadZipFile, ZipFile, ZipExtFile, + sizeFileHeader, structFileHeader, stringFileHeader, + _FH_SIGNATURE, _FH_FILENAME_LENGTH, _FH_EXTRA_FIELD_LENGTH, _FH_GENERAL_PURPOSE_FLAG_BITS, + _MASK_COMPRESSED_PATCH, _MASK_STRONG_ENCRYPTION, _MASK_UTF_FILENAME, _MASK_ENCRYPTED +) +import struct +from requests import Session + +from .webfile import WebFile + + +class _SharedWebFile(WebFile): + def __init__(self, webfile: WebFile, pos: int): + super().__init__(webfile.url, webfile.session, size=webfile.size) + self.seek(pos) + + +class WebZipFile(ZipFile): + "Lock-free version of ZipFile that reads from a WebFile, allowing for concurrent reads." + def __init__(self, url: str, session: Optional[Session] = None, headers: Optional[Dict[str, str]] = None): + """Open the ZIP file with mode read 'r', write 'w', exclusive create 'x', + or append 'a'.""" + webf = WebFile(url, session=session, headers=headers) + super().__init__(webf, mode='r') + + def open(self, name, mode="r", pwd=None, *, force_zip64=False): + """Return file-like object for 'name'. + + name is a string for the file name within the ZIP file, or a ZipInfo + object. + + mode should be 'r' to read a file already in the ZIP file, or 'w' to + write to a file newly added to the archive. + + pwd is the password to decrypt files (only used for reading). + + When writing, if the file size is not known in advance but may exceed + 2 GiB, pass force_zip64 to use the ZIP64 format, which can handle large + files. If the size is known in advance, it is best to pass a ZipInfo + instance for name, with zinfo.file_size set. + """ + if mode not in {"r", "w"}: + raise ValueError('open() requires mode "r" or "w"') + if pwd and (mode == "w"): + raise ValueError("pwd is only supported for reading files") + if not self.fp: + raise ValueError( + "Attempt to use ZIP archive that was already closed") + + assert mode == "r", "Only read mode is supported for now" + + # Make sure we have an info object + if isinstance(name, ZipInfo): + # 'name' is already an info object + zinfo = name + elif mode == 'w': + zinfo = ZipInfo(name) + zinfo.compress_type = self.compression + zinfo._compresslevel = self.compresslevel + else: + # Get info object for name + zinfo = self.getinfo(name) + + if mode == 'w': + return self._open_to_write(zinfo, force_zip64=force_zip64) + + if self._writing: + raise ValueError("Can't read from the ZIP file while there " + "is an open writing handle on it. " + "Close the writing handle before trying to read.") + + # Open for reading: + self._fileRefCnt += 1 + zef_file = _SharedWebFile(self.fp, zinfo.header_offset) + + try: + # Skip the file header: + fheader = zef_file.read(sizeFileHeader) + if len(fheader) != sizeFileHeader: + raise BadZipFile("Truncated file header") + fheader = struct.unpack(structFileHeader, fheader) + if fheader[_FH_SIGNATURE] != stringFileHeader: + raise BadZipFile("Bad magic number for file header") + + fname = zef_file.read(fheader[_FH_FILENAME_LENGTH]) + if fheader[_FH_EXTRA_FIELD_LENGTH]: + zef_file.seek(fheader[_FH_EXTRA_FIELD_LENGTH], whence=1) + + if zinfo.flag_bits & _MASK_COMPRESSED_PATCH: + # Zip 2.7: compressed patched data + raise NotImplementedError("compressed patched data (flag bit 5)") + + if zinfo.flag_bits & _MASK_STRONG_ENCRYPTION: + # strong encryption + raise NotImplementedError("strong encryption (flag bit 6)") + + if fheader[_FH_GENERAL_PURPOSE_FLAG_BITS] & _MASK_UTF_FILENAME: + # UTF-8 filename + fname_str = fname.decode("utf-8") + else: + fname_str = fname.decode(self.metadata_encoding or "cp437") + + if fname_str != zinfo.orig_filename: + raise BadZipFile( + 'File name in directory %r and header %r differ.' + % (zinfo.orig_filename, fname)) + + # check for encrypted flag & handle password + is_encrypted = zinfo.flag_bits & _MASK_ENCRYPTED + if is_encrypted: + if not pwd: + pwd = self.pwd + if pwd and not isinstance(pwd, bytes): + raise TypeError("pwd: expected bytes, got %s" % type(pwd).__name__) + if not pwd: + raise RuntimeError("File %r is encrypted, password " + "required for extraction" % name) + else: + pwd = None + + return ZipExtFile(zef_file, mode, zinfo, pwd, True) + except: + zef_file.close() + raise \ No newline at end of file diff --git a/nodes.py b/nodes.py index 356e61e..a0f279b 100644 --- a/nodes.py +++ b/nodes.py @@ -51,6 +51,101 @@ BASE_CACHE_DIR = Path(os.path.dirname(os.path.realpath(__file__))) / "triton_cac #os.environ["TRITON_ALWAYS_COMPILE"] = "1" #os.environ["TORCHINDUCTOR_FORCE_DISABLE_CACHES"]="1" +PIXAL3D_IMAGE_COND_CONFIGS = { + "ss": { + "model_name": "facebook/dinov3-vitl16-pretrain-lvd1689m", + "image_size": 512, + "grid_resolution": 16, + }, + "shape_512": { + "model_name": "facebook/dinov3-vitl16-pretrain-lvd1689m", + "image_size": 512, + "grid_resolution": 32, + "use_naf_upsample": True, + "naf_target_size": 512, + }, + "shape_1024": { + "model_name": "facebook/dinov3-vitl16-pretrain-lvd1689m", + "image_size": 1024, + "grid_resolution": 64, + "use_naf_upsample": True, + "naf_target_size": 512, + }, + "tex_1024": { + "model_name": "facebook/dinov3-vitl16-pretrain-lvd1689m", + "image_size": 1024, + "grid_resolution": 64, + "use_naf_upsample": True, + "naf_target_size": 1024, + }, +} + +def build_pixal3d_image_cond_model(config: dict): + from .trellis2.trainers.flow_matching.mixins.image_conditioned_proj import DinoV3ProjFeatureExtractor + model = DinoV3ProjFeatureExtractor(**config) + model.eval() + return model + +def load_pixal3d_image_cond_ss(pipeline, config: dict): + if hasattr(pipeline,'pixal3d_image_cond_ss') and pipeline.pixal3d_image_cond_ss is not None: + return pipeline.pixal3d_image_cond_ss + + print('Loading Pixal3D Image Cond SS Model ...') + model = build_pixal3d_image_cond_model(config) + pipeline.pixal3d_image_cond_ss = model + return model + +def unload_pixal3d_image_cond_ss(pipeline): + if hasattr(pipeline,'pixal3d_image_cond_ss') and pipeline.pixal3d_image_cond_ss is not None: + del pipeline.pixal3d_image_cond_ss + pipeline.pixal3d_image_cond_ss = None + gc.collect() + +def load_pixal3d_image_cond_shape_512(pipeline, config: dict): + if hasattr(pipeline,'pixal3d_image_cond_shape_512') and pipeline.pixal3d_image_cond_shape_512 is not None: + return pipeline.pixal3d_image_cond_shape_512 + + print('Loading Pixal3D Image Cond Shape 512 Model ...') + model = build_pixal3d_image_cond_model(config) + pipeline.pixal3d_image_cond_shape_512 = model + return model + +def unload_pixal3d_image_cond_shape_512(pipeline): + if hasattr(pipeline,'pixal3d_image_cond_shape_512') and pipeline.pixal3d_image_cond_shape_512 is not None: + del pipeline.pixal3d_image_cond_shape_512 + pipeline.pixal3d_image_cond_shape_512 = None + gc.collect() + +def load_pixal3d_image_cond_shape_1024(pipeline, config: dict): + if hasattr(pipeline,'pixal3d_image_cond_shape_1024') and pipeline.pixal3d_image_cond_shape_1024 is not None: + return pipeline.pixal3d_image_cond_shape_1024 + + print('Loading Pixal3D Image Cond Shape 1024 Model ...') + model = build_pixal3d_image_cond_model(config) + pipeline.pixal3d_image_cond_shape_1024 = model + return model + +def unload_pixal3d_image_cond_shape_1024(pipeline): + if hasattr(pipeline,'pixal3d_image_cond_shape_1024') and pipeline.pixal3d_image_cond_shape_1024 is not None: + del pipeline.pixal3d_image_cond_shape_1024 + pipeline.pixal3d_image_cond_shape_1024 = None + gc.collect() + +def load_pixal3d_image_cond_tex_1024(pipeline, config: dict): + if hasattr(pipeline,'pixal3d_image_cond_tex_1024') and pipeline.pixal3d_image_cond_tex_1024 is not None: + return pipeline.pixal3d_image_cond_tex_1024 + + print('Loading Pixal3D Image Cond Tex 1024 Model ...') + model = build_pixal3d_image_cond_model(config) + pipeline.pixal3d_image_cond_tex_1024 = model + return model + +def unload_pixal3d_image_cond_tex_1024(pipeline): + if hasattr(pipeline,'pixal3d_image_cond_tex_1024') and pipeline.pixal3d_image_cond_tex_1024 is not None: + del pipeline.pixal3d_image_cond_tex_1024 + pipeline.pixal3d_image_cond_tex_1024 = None + gc.collect() + to_pil = transforms.ToPILImage() class AnyType(str): @@ -325,7 +420,7 @@ class Trellis2LoadModel: def INPUT_TYPES(s): return { "required": { - "modelname": (["microsoft/TRELLIS.2-4B","visualbruno/TRELLIS.2-4B-FP8"],{"default":"microsoft/TRELLIS.2-4B"}), + "modelname": (["microsoft/TRELLIS.2-4B","visualbruno/TRELLIS.2-4B-FP8","TencentARC/Pixal3D-T"],{"default":"microsoft/TRELLIS.2-4B"}), "backend": (["flash_attn","xformers","sdpa","flash_attn_3"],{"default":"flash_attn"}), "device": (["cpu","cuda"],{"default":"cuda"}), "low_vram": ("BOOLEAN",{"default":True}), @@ -379,6 +474,12 @@ class Trellis2LoadModel: if not os.path.exists(dinov3_model_path): raise Exception("Facebook Dinov3 model not found in models/facebook/dinov3-vitl16-pretrain-lvd1689m folder") + dinov3_model_path = os.path.join(folder_paths.models_dir,"facebook","dinov3-vitl16-pretrain-lvd1689m") + PIXAL3D_IMAGE_COND_CONFIGS["ss"]["model_name"] = dinov3_model_path + PIXAL3D_IMAGE_COND_CONFIGS["shape_512"]["model_name"] = dinov3_model_path + PIXAL3D_IMAGE_COND_CONFIGS["shape_1024"]["model_name"] = dinov3_model_path + PIXAL3D_IMAGE_COND_CONFIGS["tex_1024"]["model_name"] = dinov3_model_path + trellis_image_large_path = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16.safetensors") if not os.path.exists(trellis_image_large_path): print('Trellis-Image-Large ss_dec_conv3d_16l8_fp16 files not found. Trying to download the files from huggingface ...') @@ -406,6 +507,9 @@ class Trellis2LoadModel: else: raise Exception("Cannot download Trellis-Image-Large file ss_dec_conv3d_16l8_fp16.safetensors") + if use_reconviagen and modelname == 'TencentARC/Pixal3D-T': + raise Exception('Model TencentARC/Pixal3D-T is not compatible with ReconViaGen') + if use_reconviagen: reconviagen_file = os.path.join(folder_paths.models_dir,'microsoft','TRELLIS.2-4B','ckpts','ss_vggt_cond.safetensors') if not os.path.exists(reconviagen_file): @@ -477,8 +581,12 @@ class Trellis2LoadModel: raise Exception("ReconViaGen cannot be used with TRELLIS.2-4B-FP8. Select microsoft/TRELLIS.2-4B") else: use_fp8 = False - - pipeline = Trellis2ImageTo3DPipeline.from_pretrained(model_path, keep_models_loaded = keep_models_loaded, use_fp8=use_fp8, use_reconviagen=use_reconviagen) + + isPixal3D = False + if modelname == "TencentARC/Pixal3D-T": + isPixal3D = True + + pipeline = Trellis2ImageTo3DPipeline.from_pretrained(model_path, keep_models_loaded = keep_models_loaded, use_fp8=use_fp8, use_reconviagen=use_reconviagen, isPixal3D = isPixal3D) pipeline.low_vram = low_vram if device=="cuda": @@ -1044,6 +1152,7 @@ class Trellis2UnWrapAndRasterizer: "use_custom_normals": ("BOOLEAN",{"default":False}), "bvh": ("BVH",), "inpainting": (["telea","ns"],{"default":"telea"}), + "reorient_vertices":(["None","90 degrees","-90 degrees"],{"default":"90 degrees"}), } } @@ -1053,7 +1162,7 @@ class Trellis2UnWrapAndRasterizer: CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True - def process(self, mesh, mesh_cluster_threshold_cone_half_angle_rad, mesh_cluster_refine_iterations, mesh_cluster_global_iterations, mesh_cluster_smooth_strength, texture_size, texture_alpha_mode, double_side_material, bake_on_vertices,use_custom_normals,bvh,inpainting): + def process(self, mesh, mesh_cluster_threshold_cone_half_angle_rad, mesh_cluster_refine_iterations, mesh_cluster_global_iterations, mesh_cluster_smooth_strength, texture_size, texture_alpha_mode, double_side_material, bake_on_vertices,use_custom_normals,bvh,inpainting,reorient_vertices): mesh_copy = copy.deepcopy(mesh) aabb = [[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]] @@ -1158,8 +1267,12 @@ class Trellis2UnWrapAndRasterizer: normals_np = out_normals.cpu().numpy() # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) - vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy() - normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy() + if reorient_vertices == '90 degrees': + vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2], -vertices_np[:, 1] + normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2], -normals_np[:, 1] + elif reorient_vertices == '-90 degrees': + vertices_np[:, 1], vertices_np[:, 2] = -vertices_np[:, 2], vertices_np[:, 1] + normals_np[:, 1], normals_np[:, 2] = -normals_np[:, 2], normals_np[:, 1] # Create mesh with vertex colors using ColorVisuals if use_custom_normals: @@ -1291,8 +1404,13 @@ class Trellis2UnWrapAndRasterizer: normals_np = out_normals.cpu().numpy() # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) - vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2], -vertices_np[:, 1] - normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2], -normals_np[:, 1] + if reorient_vertices == '90 degrees': + vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2], -vertices_np[:, 1] + normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2], -normals_np[:, 1] + elif reorient_vertices == '-90 degrees': + vertices_np[:, 1], vertices_np[:, 2] = -vertices_np[:, 2], vertices_np[:, 1] + normals_np[:, 1], normals_np[:, 2] = -normals_np[:, 2], normals_np[:, 1] + uvs_np[:, 1] = 1 - uvs_np[:, 1] # Flip UV V-coordinate if use_custom_normals: @@ -3795,13 +3913,13 @@ class Trellis2ImageCondGenerator: }, } - RETURN_TYPES = ("IMAGE_COND", "IMAGE_COND", "TRELLIS2PIPELINE",) - RETURN_NAMES = ("cond_512", "cond_1024", "pipeline",) + RETURN_TYPES = ("IMAGE_COND", "IMAGE_COND", "TRELLIS2PIPELINE", "MOGE_CAM_CONFIG") + RETURN_NAMES = ("cond_512", "cond_1024", "pipeline", "moge_camera_config") FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True - def process(self, pipeline, image, max_views,): + def process(self, pipeline, image, max_views,): images = tensor_batch_to_pil_list(image, max_views=max_views) image_in = images[0] if len(images) == 1 else images @@ -3810,15 +3928,54 @@ class Trellis2ImageCondGenerator: else: images = [image_in] - pipeline.load_image_cond_model() + pipeline.load_image_cond_model() cond_512 = pipeline.get_cond(images, 512, max_views = max_views) cond_1024 = pipeline.get_cond(images, 1024, max_views = max_views) + if pipeline.isPixal3D: + MoGeModel = "Ruicheng/moge-2-vitl" + moge_model_path = os.path.join(folder_paths.models_dir, "Ruicheng","moge-2-vitl") + + if not os.path.exists(moge_model_path): + print(f"Downloading MoGe model to: {moge_model_path}") + from huggingface_hub import snapshot_download + snapshot_download( + repo_id=MoGeModel, + local_dir=moge_model_path, + local_dir_use_symlinks=False, + ) + + moge_model_path = os.path.join(moge_model_path,'model.pt') + + print('Loading MoGe model ...') + moge_model = self.load_moge_model(model_name = moge_model_path) + + from .trellis2.utils.camera import get_camera_params_wild_moge + + if isinstance(images, list): + image = images[0] + else: + image = images + + camera_config = get_camera_params_wild_moge(image, moge_model) + print(camera_config) + + del moge_model + moge_model = None + else: + camera_config = None + if not pipeline.keep_models_loaded: pipeline.unload_image_cond_model() - return (cond_512, cond_1024, pipeline,) + return (cond_512, cond_1024, pipeline, camera_config) + + def load_moge_model(self, model_name, device='cuda'): + from .moge.model.v2 import MoGeModel + moge_model = MoGeModel.from_pretrained(model_name).to(device) + moge_model.eval() + return moge_model class Trellis2SparseGenerator: @classmethod @@ -3845,6 +4002,10 @@ class Trellis2SparseGenerator: "dino_foundation_cap": ("FLOAT",{"default":1.00,"min":0.01,"max":1.00,"step":0.01}), "keep_only_shell": ("BOOLEAN",{"default":True}), }, + "optional":{ + "image":("IMAGE",), + "moge_camera_config":("MOGE_CAM_CONFIG",) + } } RETURN_TYPES = ("COORDS", "INT", "TRELLIS2PIPELINE",) @@ -3870,9 +4031,10 @@ class Trellis2SparseGenerator: dino_substeps, hole_fill_algorithm, dino_foundation_cap, - keep_only_shell - ): - + keep_only_shell, + image = None, + moge_camera_config = None + ): self.seed_all(seed) sparse_structure_guidance_interval = [sparse_structure_guidance_interval_start,sparse_structure_guidance_interval_end] @@ -3881,7 +4043,32 @@ class Trellis2SparseGenerator: args = pipeline._pretrained_args sparse_sampler_prefix = pipeline.GetSamplerName(sparse_structure_sampler) pipeline.sparse_structure_sampler = getattr(samplers, f"Flow{sparse_sampler_prefix}GuidanceIntervalSampler")(**args['sparse_structure_sampler']['args']) - pipeline.load_sparse_structure_model() + pipeline.load_sparse_structure_model() + + if pipeline.isPixal3D: + if image is not None: + images = tensor_batch_to_pil_list(image, max_views=16) + images = list(images) + else: + raise Exception('Image is required for Pixal3D') + + if moge_camera_config is not None: + camera_angle_x = moge_camera_config['camera_angle_x'] + distance = moge_camera_config['distance'] + mesh_scale = moge_camera_config['mesh_scale'] + else: + raise Exception('MoGe Camera Config is required for Pixal3D') + + image_cond_model = load_pixal3d_image_cond_ss(pipeline, PIXAL3D_IMAGE_COND_CONFIGS["ss"]) + + image_cond = pipeline.get_proj_cond_ss( + image=images, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + image_cond_model=image_cond_model + ) + coords = pipeline.sample_sparse_structure( image_cond, sparse_structure_resolution, 1, sparse_structure_sampler_params, @@ -3896,7 +4083,8 @@ class Trellis2SparseGenerator: ) if not pipeline.keep_models_loaded: - pipeline.unload_sparse_structure_model() + pipeline.unload_sparse_structure_model() + unload_pixal3d_image_cond_ss(pipeline) return (coords, sparse_structure_resolution, pipeline,) @@ -3931,6 +4119,11 @@ class Trellis2ShapeGenerator: "dino_substeps": ("INT",{"default":4,"min":1,"max":99,"step":1}), "dino_foundation_cap": ("FLOAT",{"default":1.00,"min":0.01,"max":1.00,"step":0.01}), }, + "optional": + { + "image": ("IMAGE",), + "moge_camera_config": ("MOGE_CAM_CONFIG",), + } } RETURN_TYPES = ("SHAPE_SLAT", "INT", "TRELLIS2PIPELINE",) @@ -3951,7 +4144,9 @@ class Trellis2ShapeGenerator: verbose, dino_lock, dino_substeps, - dino_foundation_cap + dino_foundation_cap, + image = None, + moge_camera_config = None ): shape_guidance_interval = [shape_guidance_interval_start, shape_guidance_interval_end] @@ -3963,7 +4158,30 @@ class Trellis2ShapeGenerator: if resolution == 512: pipeline.unload_shape_slat_flow_model_1024() - pipeline.load_shape_slat_flow_model_512() + pipeline.load_shape_slat_flow_model_512() + + if pipeline.isPixal3D: + images = tensor_batch_to_pil_list(image, max_views=16) + image_in = images[0] if len(images) == 1 else images + + if isinstance(image_in, (list, tuple)): + images = list(image_in) + else: + images = [image_in] + + camera_angle_x = moge_camera_config['camera_angle_x'] + distance = moge_camera_config['distance'] + mesh_scale = moge_camera_config['mesh_scale'] + + image_cond_model = load_pixal3d_image_cond_shape_512(pipeline, PIXAL3D_IMAGE_COND_CONFIGS["shape_512"]) + + image_cond = pipeline.get_proj_cond_shape( + image_cond_model, images, coords, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + ) + shape_slat = pipeline.sample_shape_slat( image_cond, pipeline.models['shape_slat_flow_model_512'], coords, shape_slat_sampler_params, @@ -3975,9 +4193,34 @@ class Trellis2ShapeGenerator: if not pipeline.keep_models_loaded: pipeline.unload_shape_slat_flow_model_512() + unload_pixal3d_image_cond_shape_512(pipeline) + elif resolution == 1024: pipeline.unload_shape_slat_flow_model_512() pipeline.load_shape_slat_flow_model_1024() + + if pipeline.isPixal3D: + images = tensor_batch_to_pil_list(image, max_views=16) + image_in = images[0] if len(images) == 1 else images + + if isinstance(image_in, (list, tuple)): + images = list(image_in) + else: + images = [image_in] + + camera_angle_x = moge_camera_config['camera_angle_x'] + distance = moge_camera_config['distance'] + mesh_scale = moge_camera_config['mesh_scale'] + + image_cond_model = load_pixal3d_image_cond_shape_1024(pipeline,PIXAL3D_IMAGE_COND_CONFIGS["shape_1024"]) + + image_cond = pipeline.get_proj_cond_shape( + image_cond_model, images, coords, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + ) + shape_slat = pipeline.sample_shape_slat( image_cond, pipeline.models['shape_slat_flow_model_1024'], coords, shape_slat_sampler_params, @@ -3989,6 +4232,7 @@ class Trellis2ShapeGenerator: if not pipeline.keep_models_loaded: pipeline.unload_shape_slat_flow_model_1024() + unload_pixal3d_image_cond_shape_1024(pipeline) return (shape_slat, resolution, pipeline,) @@ -4016,6 +4260,11 @@ class Trellis2ShapeCascadeGenerator: "dino_substeps": ("INT",{"default":4,"min":1,"max":99,"step":1}), "dino_foundation_cap": ("FLOAT",{"default":1.00,"min":0.01,"max":1.00,"step":0.01}), }, + "optional": + { + "image": ("IMAGE",), + "moge_camera_config": ("MOGE_CAM_CONFIG",), + } } RETURN_TYPES = ("SHAPE_SLAT","INT","TRELLIS2PIPELINE","INT",) @@ -4036,7 +4285,9 @@ class Trellis2ShapeCascadeGenerator: verbose, dino_lock, dino_substeps, - dino_foundation_cap + dino_foundation_cap, + image = None, + moge_camera_config = None ): shape_guidance_interval = [shape_guidance_interval_start, shape_guidance_interval_end] @@ -4046,14 +4297,14 @@ class Trellis2ShapeCascadeGenerator: shape_sampler_prefix = pipeline.GetSamplerName(shape_sampler) pipeline.shape_slat_sampler = getattr(samplers, f"Flow{shape_sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args']) pipeline.load_shape_slat_flow_model_1024() - slat, hr_resolution, num_tokens = self.sample(pipeline, shape_slat, from_resolution, to_resolution, sparse_structure_resolution, max_num_tokens, image_cond, shape_slat_sampler_params, pipeline.models['shape_slat_flow_model_1024'], verbose, dino_lock, dino_substeps, dino_foundation_cap) + slat, hr_resolution, num_tokens = self.sample(pipeline, shape_slat, from_resolution, to_resolution, sparse_structure_resolution, max_num_tokens, image_cond, shape_slat_sampler_params, pipeline.models['shape_slat_flow_model_1024'], verbose, dino_lock, dino_substeps, dino_foundation_cap, image, moge_camera_config) if not pipeline.keep_models_loaded: pipeline.unload_shape_slat_flow_model_1024() return (slat, hr_resolution, pipeline, num_tokens,) - def sample(self, pipeline, slat, lr_resolution, resolution, sparse_structure_resolution, max_num_tokens, cond, sampler_params, flow_model, verbose, dino_lock, dino_substeps, dino_foundation_cap): + def sample(self, pipeline, slat, lr_resolution, resolution, sparse_structure_resolution, max_num_tokens, cond, sampler_params, flow_model, verbose, dino_lock, dino_substeps, dino_foundation_cap, image, moge_camera_config): # Upsample pipeline.load_shape_slat_decoder() if pipeline.low_vram: @@ -4092,6 +4343,28 @@ class Trellis2ShapeCascadeGenerator: hr_resolution = 512 break + if pipeline.isPixal3D: + images = tensor_batch_to_pil_list(image, max_views=16) + image_in = images[0] if len(images) == 1 else images + + if isinstance(image_in, (list, tuple)): + images = list(image_in) + else: + images = [image_in] + + camera_angle_x = moge_camera_config['camera_angle_x'] + distance = moge_camera_config['distance'] + mesh_scale = moge_camera_config['mesh_scale'] + + image_cond_model = load_pixal3d_image_cond_shape_1024(pipeline, PIXAL3D_IMAGE_COND_CONFIGS["shape_1024"]) + + cond = pipeline.get_proj_cond_shape( + image_cond_model, images, coords, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + ) + if pipeline.low_vram: cond = pipeline._cond_to(cond, pipeline.device) @@ -4127,6 +4400,7 @@ class Trellis2ShapeCascadeGenerator: if pipeline.low_vram: cond = pipeline._cond_cpu(cond) pipeline._cleanup_cuda() + unload_pixal3d_image_cond_shape_1024(pipeline) return slat, hr_resolution, num_tokens @@ -4151,6 +4425,11 @@ class Trellis2TexSlatGenerator: "dino_substeps": ("INT",{"default":4,"min":1,"max":99,"step":1}), "dino_foundation_cap": ("FLOAT",{"default":1.00,"min":0.01,"max":1.00,"step":0.01}), }, + "optional": + { + "image": ("IMAGE",), + "moge_camera_config": ("MOGE_CAM_CONFIG",), + } } RETURN_TYPES = ("TEXTURE_SLAT", "TRELLIS2PIPELINE",) @@ -4171,15 +4450,21 @@ class Trellis2TexSlatGenerator: verbose, dino_lock, dino_substeps, - dino_foundation_cap + dino_foundation_cap, + image = None, + moge_camera_config = None ): texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end] tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"guidance_interval":texture_guidance_interval,"rescale_t":texture_rescale_t} if resolution == 512: + if pipeline.isPixal3D: + raise Exception('Pixal3D only works with 1024 resolution') + pipeline.unload_tex_slat_flow_model_1024() - pipeline.load_tex_slat_flow_model_512() + pipeline.load_tex_slat_flow_model_512() + tex_slat = pipeline.sample_tex_slat_advanced( image_cond, pipeline.models['tex_slat_flow_model_512'], shape_slat, tex_slat_sampler_params, texture_sampler, @@ -4192,8 +4477,36 @@ class Trellis2TexSlatGenerator: pipeline.unload_tex_slat_flow_model_512() elif resolution == 1024: - pipeline.unload_tex_slat_flow_model_512() + if not pipeline.isPixal3D: + pipeline.unload_tex_slat_flow_model_512() + pipeline.load_tex_slat_flow_model_1024() + + if pipeline.isPixal3D: + images = tensor_batch_to_pil_list(image, max_views=16) + image_in = images[0] if len(images) == 1 else images + + if isinstance(image_in, (list, tuple)): + images = list(image_in) + else: + images = [image_in] + + camera_angle_x = moge_camera_config['camera_angle_x'] + distance = moge_camera_config['distance'] + mesh_scale = moge_camera_config['mesh_scale'] + + image_cond_model = load_pixal3d_image_cond_tex_1024(pipeline, PIXAL3D_IMAGE_COND_CONFIGS["tex_1024"]) + + tex_grid_res = resolution // 16 + + image_cond = pipeline.get_proj_cond_shape( + image_cond_model, images, shape_slat.coords, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + grid_resolution_override=tex_grid_res, + ) + tex_slat = pipeline.sample_tex_slat_advanced( image_cond, pipeline.models['tex_slat_flow_model_1024'], shape_slat, tex_slat_sampler_params, texture_sampler, @@ -4365,7 +4678,7 @@ class Trellis2MultiViewTexturing: "trimesh": ("TRIMESH",), "texture_size": ("INT", {"default": 4096, "min": 512, "max": 8192}), "blend_texture": ("BOOLEAN", {"default":True}), - "blend_exponent": ("FLOAT", {"default": 1.0, "min": 0.5, "max": 8.0, "step": 0.5}), + "blend_exponent": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 99.0, "step": 0.1}), "ortho_scale": ("FLOAT", {"default": 1.1, "min": 0.05, "max": 10.0, "step": 0.01}), "norm_size": ("FLOAT",{"default":1.15, "min":0.0, "max":9.99, "step":0.01}), "fill_holes": ("BOOLEAN",{"default":True}), @@ -4373,6 +4686,7 @@ class Trellis2MultiViewTexturing: "use_metallic": ("BOOLEAN",{"default":True}), "depth_eps": ("FLOAT",{"default":0.0100,"min":0.0001,"max":1.0000,"step":0.0001}), "mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60,"min":1,"max":179,"step":1}), + "add_alpha_channel": ("BOOLEAN",{"default":False}), }, "optional": { # Standard views @@ -4416,6 +4730,7 @@ class Trellis2MultiViewTexturing: use_metallic, depth_eps, mesh_cluster_threshold_cone_half_angle_rad, + add_alpha_channel, baseColorTexture = None, front_image=None, back_image=None, @@ -4513,7 +4828,8 @@ class Trellis2MultiViewTexturing: norm_size=norm_size, max_hole_size=max_hole_size, use_metallic=use_metallic, - depth_eps=depth_eps + depth_eps=depth_eps, + add_alpha_channel=add_alpha_channel ) return (trimesh_obj, pil2tensor(base_color), pil2tensor(mr)) @@ -4563,13 +4879,15 @@ class Trellis2ProjectHighPolyToLowPoly: "low_poly_trimesh": ("TRIMESH",), "texture_size": ("INT", {"default": 4096, "min": 512, "max": 8192}), "blend_texture": ("BOOLEAN", {"default":True}), - "blend_exponent": ("FLOAT", {"default": 1.0, "min": 0.5, "max": 8.0, "step": 0.5}), + "blend_exponent": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 99.0, "step": 0.5}), "ortho_scale": ("FLOAT", {"default": 1.1, "min": 0.05, "max": 10.0, "step": 0.01}), "norm_size": ("FLOAT",{"default":1.15, "min":0.0, "max":9.99, "step":0.01}), "fill_holes": ("BOOLEAN",{"default":True}), "max_hole_size": ("INT",{"default":20,"min":0,"max":99999,"step":1}), - "use_metallic": ("BOOLEAN",{"default":True}), + "use_metallic": ("BOOLEAN",{"default":False}), "depth_eps": ("FLOAT",{"default":0.0100,"min":0.0001,"max":1.0000,"step":0.0001}), + "mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60,"min":1,"max":179,"step":1}), + "add_alpha_channel": ("BOOLEAN",{"default":False}), }, "optional": { # Standard views @@ -4613,6 +4931,8 @@ class Trellis2ProjectHighPolyToLowPoly: max_hole_size, use_metallic, depth_eps, + mesh_cluster_threshold_cone_half_angle_rad, + add_alpha_channel, baseColorTexture = None, front_image=None, back_image=None, @@ -4702,6 +5022,7 @@ class Trellis2ProjectHighPolyToLowPoly: elevations, weights, texture_size=texture_size, + mesh_cluster_threshold_cone_half_angle_rad=mesh_cluster_threshold_cone_half_angle_rad, blend_exponent=blend_exponent, ortho_scale=ortho_scale, blend_texture=blend_texture, @@ -4710,7 +5031,8 @@ class Trellis2ProjectHighPolyToLowPoly: max_hole_size=max_hole_size, use_metallic=use_metallic, depth_eps=depth_eps, - low_poly_mesh=low_poly_trimesh + low_poly_mesh=low_poly_trimesh, + add_alpha_channel=add_alpha_channel ) return (trimesh_obj, pil2tensor(base_color), pil2tensor(mr)) @@ -6419,7 +6741,7 @@ class Trellis2TexSlatMultiViewGenerator: conds[v] = pipeline._cond_cpu(conds[v]) pipeline._cleanup_cuda() - return slat + return slat NODE_CLASS_MAPPINGS = { "Trellis2LoadModel": Trellis2LoadModel, diff --git a/trellis2/models/sparse_structure_flow.py b/trellis2/models/sparse_structure_flow.py index 60baf1c..6e04de1 100644 --- a/trellis2/models/sparse_structure_flow.py +++ b/trellis2/models/sparse_structure_flow.py @@ -73,6 +73,9 @@ class SparseStructureFlowModel(nn.Module): initialization: str = 'vanilla', qk_rms_norm: bool = False, qk_rms_norm_cross: bool = False, + image_attn_mode: Literal["cross", "proj", "gated_proj"] = "cross", + proj_in_channels: Optional[int] = None, + vae_in_channels: Optional[int] = None, **kwargs ): super().__init__() @@ -90,6 +93,9 @@ class SparseStructureFlowModel(nn.Module): self.initialization = initialization self.qk_rms_norm = qk_rms_norm self.qk_rms_norm_cross = qk_rms_norm_cross + self.image_attn_mode = image_attn_mode + self.proj_in_channels = proj_in_channels + self.vae_in_channels = vae_in_channels self.dtype = str_to_dtype(dtype) self.t_embedder = TimestepEmbedder(model_channels) @@ -130,6 +136,9 @@ class SparseStructureFlowModel(nn.Module): share_mod=share_mod, qk_rms_norm=self.qk_rms_norm, qk_rms_norm_cross=self.qk_rms_norm_cross, + image_attn_mode=image_attn_mode, + proj_in_channels=proj_in_channels, + vae_in_channels=vae_in_channels, ) for _ in range(num_blocks) ]) @@ -199,7 +208,11 @@ class SparseStructureFlowModel(nn.Module): nn.init.constant_(module.bias, 0) for block in self.blocks: block.self_attn.to_out.apply(_scaled_init) - block.cross_attn.to_out.apply(_scaled_init) + # Handle cross, proj, and gated_proj modes + if self.image_attn_mode in ("proj", "gated_proj"): + block.cross_attn.cross_attn_block.to_out.apply(_scaled_init) + else: + block.cross_attn.to_out.apply(_scaled_init) block.mlp.mlp[2].apply(_scaled_init) # Initialize input layer to make the initial representation have variance 1 @@ -224,6 +237,19 @@ class SparseStructureFlowModel(nn.Module): nn.init.constant_(self.out_layer.bias, 0) def forward(self, x: torch.Tensor, t: torch.Tensor, cond: torch.Tensor) -> torch.Tensor: + """ + Forward pass. + + Args: + x: Input tensor [B, C, D, H, W] + t: Timestep tensor [B] + cond: Conditioning tensor. For "cross" mode: [B, N, D]. + For "proj" mode: dict {'global': global_cond, 'proj': proj_cond} + or tuple of (global_cond, proj_cond) + + Returns: + Output tensor [B, C, D, H, W] + """ assert [*x.shape] == [x.shape[0], self.in_channels, *[self.resolution] * 3], \ f"Input shape mismatch, got {x.shape}, expected {[x.shape[0], self.in_channels, *[self.resolution] * 3]}" @@ -237,7 +263,28 @@ class SparseStructureFlowModel(nn.Module): t_emb = self.adaLN_modulation(t_emb) t_emb = manual_cast(t_emb, self.dtype) h = manual_cast(h, self.dtype) - cond = manual_cast(cond, self.dtype) + + # Handle different conditioning modes + if hasattr(self,'image_attn_mode'): + if self.image_attn_mode == 'proj': + if isinstance(cond, dict): + global_cond = cond['global'] + proj_cond = cond['proj'] + else: + global_cond, proj_cond = cond + global_cond = manual_cast(global_cond, self.dtype) + proj_cond = manual_cast(proj_cond, self.dtype) + cond = (global_cond, proj_cond) + elif self.image_attn_mode == 'gated_proj': + global_cond = manual_cast(cond['global'], self.dtype) + proj_semantic = manual_cast(cond['proj_semantic'], self.dtype) + proj_color = manual_cast(cond['proj_color'], self.dtype) + cond = {'global': global_cond, 'proj_semantic': proj_semantic, 'proj_color': proj_color} + else: + cond = manual_cast(cond, self.dtype) + else: + cond = manual_cast(cond, self.dtype) + for block in self.blocks: h = block(h, t_emb, cond, self.rope_phases) h = manual_cast(h, x.dtype) diff --git a/trellis2/models/structured_latent_flow.py b/trellis2/models/structured_latent_flow.py index 595ec5a..bf33868 100644 --- a/trellis2/models/structured_latent_flow.py +++ b/trellis2/models/structured_latent_flow.py @@ -13,6 +13,13 @@ from .sparse_elastic_mixin import SparseTransformerElasticMixin class SLatFlowModel(nn.Module): + """ + Structured Latent Flow Model for 3D generation. + + Supports two conditioning modes: + - "cross": Standard cross-attention with image features + - "proj": View-aligned projection attention with camera-aware features + """ def __init__( self, resolution: int, @@ -32,6 +39,10 @@ class SLatFlowModel(nn.Module): initialization: str = 'vanilla', qk_rms_norm: bool = False, qk_rms_norm_cross: bool = False, + image_attn_mode: Literal["cross", "proj", "gated_proj"] = "cross", + proj_in_channels: Optional[int] = None, + vae_in_channels: Optional[int] = None, + **kwargs ): super().__init__() self.resolution = resolution @@ -48,6 +59,9 @@ class SLatFlowModel(nn.Module): self.initialization = initialization self.qk_rms_norm = qk_rms_norm self.qk_rms_norm_cross = qk_rms_norm_cross + self.image_attn_mode = image_attn_mode + self.proj_in_channels = proj_in_channels + self.vae_in_channels = vae_in_channels self.dtype = str_to_dtype(dtype) self.t_embedder = TimestepEmbedder(model_channels) @@ -75,6 +89,9 @@ class SLatFlowModel(nn.Module): share_mod=self.share_mod, qk_rms_norm=self.qk_rms_norm, qk_rms_norm_cross=self.qk_rms_norm_cross, + image_attn_mode=image_attn_mode, + proj_in_channels=proj_in_channels, + vae_in_channels=vae_in_channels, ) for _ in range(num_blocks) ]) @@ -144,7 +161,11 @@ class SLatFlowModel(nn.Module): nn.init.constant_(module.bias, 0) for block in self.blocks: block.self_attn.to_out.apply(_scaled_init) - block.cross_attn.to_out.apply(_scaled_init) + # Handle cross, proj, and gated_proj modes + if self.image_attn_mode in ("proj", "gated_proj"): + block.cross_attn.cross_attn_block.to_out.apply(_scaled_init) + else: + block.cross_attn.to_out.apply(_scaled_init) block.mlp.mlp[2].apply(_scaled_init) # Initialize input layer to make the initial representation have variance 1 @@ -172,14 +193,26 @@ class SLatFlowModel(nn.Module): self, x: sp.SparseTensor, t: torch.Tensor, - cond: Union[torch.Tensor, List[torch.Tensor]], + cond: Union[torch.Tensor, List[torch.Tensor], Dict[str, Union[torch.Tensor, sp.SparseTensor]], Tuple], concat_cond: Optional[sp.SparseTensor] = None, **kwargs ) -> sp.SparseTensor: + """ + Forward pass. + + Args: + x: SparseTensor input + t: Timestep tensor [B] + cond: Conditioning tensor. For "cross" mode: list of tensors or tensor. + For "proj" mode: dict {'global': global_cond, 'proj': proj_cond} + or tuple of (global_cond, proj_cond) + concat_cond: Optional concatenation condition + + Returns: + SparseTensor output + """ if concat_cond is not None: x = sp.sparse_cat([x, concat_cond], dim=-1) - if isinstance(cond, list): - cond = sp.VarLenTensor.from_tensor_list(cond) h = self.input_layer(x) h = manual_cast(h, self.dtype) @@ -187,11 +220,36 @@ class SLatFlowModel(nn.Module): if self.share_mod: t_emb = self.adaLN_modulation(t_emb) t_emb = manual_cast(t_emb, self.dtype) - cond = manual_cast(cond, self.dtype) if self.pe_mode == "ape": pe = self.pos_embedder(h.coords[:, 1:]) h = h + manual_cast(pe, self.dtype) + + # Handle different conditioning modes + if self.image_attn_mode == 'proj': + if isinstance(cond, dict): + global_cond = cond['global'] + proj_cond = cond['proj'] + else: + global_cond, proj_cond = cond + if isinstance(global_cond, list): + global_cond = sp.VarLenTensor.from_tensor_list(global_cond) + global_cond = manual_cast(global_cond, self.dtype) + proj_cond = manual_cast(proj_cond, self.dtype) + cond = (global_cond, proj_cond) + elif self.image_attn_mode == 'gated_proj': + global_cond = cond['global'] + if isinstance(global_cond, list): + global_cond = sp.VarLenTensor.from_tensor_list(global_cond) + global_cond = manual_cast(global_cond, self.dtype) + proj_semantic = manual_cast(cond['proj_semantic'], self.dtype) + proj_color = manual_cast(cond['proj_color'], self.dtype) + cond = {'global': global_cond, 'proj_semantic': proj_semantic, 'proj_color': proj_color} + else: + if isinstance(cond, list): + cond = sp.VarLenTensor.from_tensor_list(cond) + cond = manual_cast(cond, self.dtype) + for block in self.blocks: h = block(h, t_emb, cond) diff --git a/trellis2/modules/attention/__init__.py b/trellis2/modules/attention/__init__.py index e90e901..95a7e79 100644 --- a/trellis2/modules/attention/__init__.py +++ b/trellis2/modules/attention/__init__.py @@ -1,3 +1,4 @@ from .full_attn import * from .modules import * from .rope import * +from .proj_attention import ProjectAttention, GatedProjectAttention \ No newline at end of file diff --git a/trellis2/modules/attention/proj_attention.py b/trellis2/modules/attention/proj_attention.py new file mode 100644 index 0000000..77011c1 --- /dev/null +++ b/trellis2/modules/attention/proj_attention.py @@ -0,0 +1,101 @@ +""" +View-Aligned Projection Attention Module for TRELLIS2 + +This module implements the projection-based attention mechanism that combines +global cross-attention with view-aligned projected features. + +Supports two modes: +- "proj": Standard projection (DINOv3 only), per-block proj_linear +- "gated_proj": Gated fusion of DINOv3 (semantic) + VAE (color) features +""" + +from typing import * +import torch +import torch.nn as nn + + +class ProjectAttention(nn.Module): + """ + Projection-based Attention Module with per-block proj_linear. + + Combines global cross-attention with view-aligned projected features. + Each block owns a proj_linear that projects DINOv3 features from + proj_in_channels (e.g. 1024) to model_channels (e.g. 1536). + + The module receives: + - x: Input features from the transformer + - context: A dict with keys: + - 'global': Global image features, shape [B, M, ctx_channels] + - 'proj': View-aligned projected features, shape [B, N, proj_in_channels] + + The output combines the cross-attention result with the projected context. + """ + def __init__(self, cross_attn_block: nn.Module, channels: int, proj_in_channels: int): + super().__init__() + self.cross_attn_block = cross_attn_block + self.proj_linear = nn.Linear(proj_in_channels, channels, bias=True) + + def forward(self, x: torch.Tensor, context: Union[Dict[str, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]) -> torch.Tensor: + if isinstance(context, dict): + global_context = context['global'] + proj_context = context['proj'] + else: + global_context, proj_context = context + + global_out = self.cross_attn_block(x, global_context) + proj_out = self.proj_linear(proj_context) + context_combined = proj_out + global_out + return context_combined + + +class GatedProjectAttention(nn.Module): + """ + Concat-Projection Attention Module for DINOv3 (semantic) + VAE (color) features. + + Concatenates DINOv3 and VAE projected features and applies a single linear + projection to model_channels. This is mathematically equivalent to two + separate proj_linears + addition, but allows cross-dimensional interactions + between semantic and color features through the shared weight matrix. + + Zero-initialized for stable training: at init, fused=0 so only global + cross-attention contributes; color+semantic signals are gradually learned. + + The module receives: + - x: Input features from the transformer + - context: A dict with keys: + - 'global': Global image features, shape [B, M, ctx_channels] + - 'proj_semantic': DINOv3 projected features, shape [B, N, dino_channels] + - 'proj_color': VAE projected features, shape [B, N, vae_channels] + """ + def __init__( + self, + cross_attn_block: nn.Module, + channels: int, + dino_in_channels: int, + vae_in_channels: int, + ): + """ + Args: + cross_attn_block: The underlying cross-attention module + channels: Model channels (output dimension) + dino_in_channels: DINOv3 proj feature dimension (e.g. 1024) + vae_in_channels: VAE latent feature dimension (e.g. 16) + """ + super().__init__() + self.cross_attn_block = cross_attn_block + self.proj_linear = nn.Linear(dino_in_channels + vae_in_channels, channels, bias=True) + # Zero-init: at start, fused=0, only global cross-attn contributes + nn.init.zeros_(self.proj_linear.weight) + nn.init.zeros_(self.proj_linear.bias) + + def forward(self, x: torch.Tensor, context: Union[Dict[str, torch.Tensor], Tuple]) -> torch.Tensor: + if isinstance(context, dict): + global_context = context['global'] + proj_semantic = context['proj_semantic'] + proj_color = context['proj_color'] + else: + global_context, proj_semantic, proj_color = context + + global_out = self.cross_attn_block(x, global_context) + fused = self.proj_linear(torch.cat([proj_semantic, proj_color], dim=-1)) + return fused + global_out diff --git a/trellis2/modules/sparse/attention/__init__.py b/trellis2/modules/sparse/attention/__init__.py index 18ab3cc..c297c9c 100644 --- a/trellis2/modules/sparse/attention/__init__.py +++ b/trellis2/modules/sparse/attention/__init__.py @@ -1,3 +1,4 @@ from .full_attn import * from .windowed_attn import * from .modules import * +from .proj_attention import * \ No newline at end of file diff --git a/trellis2/modules/sparse/attention/proj_attention.py b/trellis2/modules/sparse/attention/proj_attention.py new file mode 100644 index 0000000..185c1db --- /dev/null +++ b/trellis2/modules/sparse/attention/proj_attention.py @@ -0,0 +1,99 @@ +""" +Sparse View-Aligned Projection Attention Module for TRELLIS2 + +Sparse versions of ProjectAttention and GatedProjectAttention. + +Supports two modes: +- "proj": Standard projection (DINOv3 only) +- "gated_proj": Gated fusion of DINOv3 (semantic) + VAE (color) features +""" + +from typing import * +import torch +import torch.nn as nn +from ..basic import SparseTensor, VarLenTensor + + +class SparseProjectAttention(nn.Module): + """ + Sparse Projection-based Attention Module with per-block proj_linear. + """ + def __init__(self, cross_attn_block: nn.Module, channels: int, proj_in_channels: int): + super().__init__() + self.cross_attn_block = cross_attn_block + self.proj_linear = nn.Linear(proj_in_channels, channels, bias=True) + + def forward( + self, + x: SparseTensor, + context: Union[Dict[str, Union[torch.Tensor, VarLenTensor, SparseTensor]], + Tuple[Union[torch.Tensor, VarLenTensor], SparseTensor]] + ) -> SparseTensor: + if isinstance(context, dict): + global_context = context['global'] + proj_context = context['proj'] + else: + global_context, proj_context = context + + global_out = self.cross_attn_block(x, global_context) + + if isinstance(proj_context, SparseTensor): + proj_feats = self.proj_linear(proj_context.feats) + combined_feats = proj_feats + global_out.feats + else: + proj_feats = self.proj_linear(proj_context) + combined_feats = proj_feats + global_out.feats + + return global_out.replace(combined_feats) + + +class SparseGatedProjectAttention(nn.Module): + """ + Sparse Concat-Projection Attention Module for DINOv3 + VAE features. + + Concatenates DINOv3 and VAE projected features and applies a single linear + projection to model_channels. Zero-initialized for stable training. + + Context dict must contain: + - 'global': Global image features for cross-attention + - 'proj_semantic': DINOv3 projected features (SparseTensor or Tensor) + - 'proj_color': VAE projected features (SparseTensor or Tensor) + """ + def __init__( + self, + cross_attn_block: nn.Module, + channels: int, + dino_in_channels: int, + vae_in_channels: int, + ): + super().__init__() + self.cross_attn_block = cross_attn_block + self.proj_linear = nn.Linear(dino_in_channels + vae_in_channels, channels, bias=True) + # Zero-init: at start, fused=0, only global cross-attn contributes + nn.init.zeros_(self.proj_linear.weight) + nn.init.zeros_(self.proj_linear.bias) + + def _get_feats(self, t): + return t.feats if isinstance(t, SparseTensor) else t + + def forward( + self, + x: SparseTensor, + context: Union[Dict[str, Union[torch.Tensor, VarLenTensor, SparseTensor]], Tuple], + ) -> SparseTensor: + if isinstance(context, dict): + global_context = context['global'] + proj_semantic = context['proj_semantic'] + proj_color = context['proj_color'] + else: + global_context, proj_semantic, proj_color = context + + global_out = self.cross_attn_block(x, global_context) + + fused = self.proj_linear(torch.cat([ + self._get_feats(proj_semantic), + self._get_feats(proj_color), + ], dim=-1)) + combined_feats = fused + global_out.feats + + return global_out.replace(combined_feats) diff --git a/trellis2/modules/sparse/transformer/modulated.py b/trellis2/modules/sparse/transformer/modulated.py index ad74a9c..9f5e33d 100644 --- a/trellis2/modules/sparse/transformer/modulated.py +++ b/trellis2/modules/sparse/transformer/modulated.py @@ -2,7 +2,7 @@ from typing import * import torch import torch.nn as nn from ..basic import VarLenTensor, SparseTensor -from ..attention import SparseMultiHeadAttention +from ..attention import SparseMultiHeadAttention, SparseProjectAttention, SparseGatedProjectAttention from ...norm import LayerNorm32 from .blocks import SparseFeedForwardNet @@ -81,6 +81,10 @@ class ModulatedSparseTransformerBlock(nn.Module): class ModulatedSparseTransformerCrossBlock(nn.Module): """ Sparse Transformer cross-attention block (MSA + MCA + FFN) with adaptive layer norm conditioning. + + Supports two image attention modes: + - "cross": Standard cross-attention with image features + - "proj": Projection-based attention with view-aligned features """ def __init__( self, @@ -98,11 +102,14 @@ class ModulatedSparseTransformerCrossBlock(nn.Module): qk_rms_norm_cross: bool = False, qkv_bias: bool = True, share_mod: bool = False, - + image_attn_mode: Literal["cross", "proj", "gated_proj"] = "cross", + proj_in_channels: Optional[int] = None, + vae_in_channels: Optional[int] = None, ): super().__init__() self.use_checkpoint = use_checkpoint self.share_mod = share_mod + self.image_attn_mode = image_attn_mode self.norm1 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6) self.norm2 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6) self.norm3 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6) @@ -118,15 +125,54 @@ class ModulatedSparseTransformerCrossBlock(nn.Module): rope_freq=rope_freq, qk_rms_norm=qk_rms_norm, ) - self.cross_attn = SparseMultiHeadAttention( - channels, - ctx_channels=ctx_channels, - num_heads=num_heads, - type="cross", - attn_mode="full", - qkv_bias=qkv_bias, - qk_rms_norm=qk_rms_norm_cross, - ) + + # Build cross attention based on mode + if image_attn_mode == "cross": + self.cross_attn = SparseMultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + elif image_attn_mode == "proj": + _proj_in = proj_in_channels if proj_in_channels is not None else ctx_channels + cross_attn_block = SparseMultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + self.cross_attn = SparseProjectAttention(cross_attn_block, channels, _proj_in) + elif image_attn_mode == "gated_proj": + _dino_in = proj_in_channels if proj_in_channels is not None else ctx_channels + _vae_in = vae_in_channels if vae_in_channels is not None else 16 + cross_attn_block = SparseMultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + self.cross_attn = SparseGatedProjectAttention(cross_attn_block, channels, _dino_in, _vae_in) + else: + self.cross_attn = SparseMultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + self.mlp = SparseFeedForwardNet( channels, mlp_ratio=mlp_ratio, diff --git a/trellis2/modules/transformer/modulated.py b/trellis2/modules/transformer/modulated.py index e2f6923..be029d0 100644 --- a/trellis2/modules/transformer/modulated.py +++ b/trellis2/modules/transformer/modulated.py @@ -1,7 +1,7 @@ from typing import * import torch import torch.nn as nn -from ..attention import MultiHeadAttention +from ..attention import MultiHeadAttention, ProjectAttention, GatedProjectAttention from ..norm import LayerNorm32 from .blocks import FeedForwardNet @@ -96,6 +96,9 @@ class ModulatedTransformerCrossBlock(nn.Module): qk_rms_norm_cross: bool = False, qkv_bias: bool = True, share_mod: bool = False, + image_attn_mode: Literal["cross", "proj", "gated_proj"] = "cross", + proj_in_channels: Optional[int] = None, + vae_in_channels: Optional[int] = None, ): super().__init__() self.use_checkpoint = use_checkpoint @@ -115,15 +118,46 @@ class ModulatedTransformerCrossBlock(nn.Module): rope_freq=rope_freq, qk_rms_norm=qk_rms_norm, ) - self.cross_attn = MultiHeadAttention( - channels, - ctx_channels=ctx_channels, - num_heads=num_heads, - type="cross", - attn_mode="full", - qkv_bias=qkv_bias, - qk_rms_norm=qk_rms_norm_cross, - ) + + # Build cross attention based on mode + if image_attn_mode == "cross": + self.cross_attn = MultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + elif image_attn_mode == "proj": + _proj_in = proj_in_channels if proj_in_channels is not None else ctx_channels + cross_attn_block = MultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + self.cross_attn = ProjectAttention(cross_attn_block, channels, _proj_in) + elif image_attn_mode == "gated_proj": + _dino_in = proj_in_channels if proj_in_channels is not None else ctx_channels + _vae_in = vae_in_channels if vae_in_channels is not None else 16 + cross_attn_block = MultiHeadAttention( + channels, + ctx_channels=ctx_channels, + num_heads=num_heads, + type="cross", + attn_mode="full", + qkv_bias=qkv_bias, + qk_rms_norm=qk_rms_norm_cross, + ) + self.cross_attn = GatedProjectAttention(cross_attn_block, channels, _dino_in, _vae_in) + else: + raise ValueError(f"Unknown image attention mode: {image_attn_mode}") + self.mlp = FeedForwardNet( channels, mlp_ratio=mlp_ratio, diff --git a/trellis2/pipelines/samplers/flow_euler.py b/trellis2/pipelines/samplers/flow_euler.py index 3d44733..22d8fbf 100644 --- a/trellis2/pipelines/samplers/flow_euler.py +++ b/trellis2/pipelines/samplers/flow_euler.py @@ -76,7 +76,7 @@ class DinoLockMixin: def _dino_lock_step(self, model, x_t, t, t_prev, cond, lock_strength, step_idx, total_steps, - substeps=1, v_ema=None, ema_alpha=0.8, verbose=True, **kwargs): + substeps=1, v_ema=None, ema_alpha=0.85, verbose=True, **kwargs): """ One step with DINO lock + velocity EMA smoothing. diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index c422771..39ffbf0 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -151,7 +151,7 @@ class Trellis2ImageTo3DPipeline(Pipeline): torch.cuda.empty_cache() @classmethod - def from_pretrained(cls, path: str, config_file: str = "pipeline.json", keep_models_loaded = True, use_fp8 = False, use_reconviagen = False) -> "Trellis2ImageTo3DPipeline": + def from_pretrained(cls, path: str, config_file: str = "pipeline.json", keep_models_loaded = True, use_fp8 = False, use_reconviagen = False, isPixal3D = False) -> "Trellis2ImageTo3DPipeline": """ Load a pretrained model. @@ -197,8 +197,11 @@ class Trellis2ImageTo3DPipeline(Pipeline): pipeline.keep_models_loaded = keep_models_loaded pipeline.last_processing = '' pipeline.use_fp8 = use_fp8 + pipeline.isPixal3D = isPixal3D - pipeline._pretrained_args['models']['sparse_structure_decoder'] = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16") + if not isPixal3D: + pipeline._pretrained_args['models']['sparse_structure_decoder'] = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16") + facebook_model_path = os.path.join(folder_paths.models_dir,"facebook","dinov3-vitl16-pretrain-lvd1689m") pipeline._pretrained_args['image_cond_model']['args']['model_name'] = facebook_model_path @@ -211,8 +214,11 @@ class Trellis2ImageTo3DPipeline(Pipeline): self.models['sparse_structure_flow_model'].eval() self.models['sparse_structure_flow_model'].to(self._device) - if self.models['sparse_structure_decoder'] is None: - self.models['sparse_structure_decoder'] = models.from_pretrained(self._pretrained_args['models']['sparse_structure_decoder']) + if self.models['sparse_structure_decoder'] is None: + if self.isPixal3D: + self.models['sparse_structure_decoder'] = models.from_pretrained(f"{self.path}/{self._pretrained_args['models']['sparse_structure_decoder']}") + else: + self.models['sparse_structure_decoder'] = models.from_pretrained(self._pretrained_args['models']['sparse_structure_decoder']) self.models['sparse_structure_decoder'].eval() self.models['sparse_structure_decoder'].to(self._device) if hasattr(self.models['sparse_structure_decoder'], 'low_vram'): @@ -389,6 +395,118 @@ class Trellis2ImageTo3DPipeline(Pipeline): if self.rembg_model is not None: self.rembg_model.to(device) + # ========================================================================= + # Proj mode condition building + # ========================================================================= + + @torch.no_grad() + def get_proj_cond_ss( + self, + image: list, + camera_angle_x: float = 0.8575560450553894, + distance: float = 2.0, + mesh_scale: float = 1.0, + image_cond_model = None + ) -> dict: + """ + Get proj conditioning for sparse structure stage. + + Args: + image: List of PIL images. + camera_angle_x: Camera horizontal FOV in radians. + distance: Camera distance. + mesh_scale: Mesh scale. + + Returns: + dict with 'cond' and 'neg_cond', each containing {'global': ..., 'proj': ...} + """ + print('Getting Proj Image Cond ...') + device = self.device + #image_cond_model = self.image_cond_model + if self.low_vram: + image_cond_model.to(device) + cam_angle = torch.tensor([camera_angle_x], device=device) + dist_tensor = torch.tensor([distance], device=device) + scale_tensor = torch.tensor([mesh_scale], device=device) + z_global, z_proj = image_cond_model( + image, camera_angle_x=cam_angle, distance=dist_tensor, mesh_scale=scale_tensor, + ) + if self.low_vram: + image_cond_model.cpu() + return { + 'cond': {'global': z_global, 'proj': z_proj}, + 'neg_cond': {'global': torch.zeros_like(z_global), 'proj': torch.zeros_like(z_proj)}, + } + + @torch.no_grad() + def get_proj_cond_shape( + self, + image_cond_model: nn.Module, + image: list, + coords: torch.Tensor, + camera_angle_x: float = 0.8575560450553894, + distance: float = 2.0, + mesh_scale: float = 1.0, + grid_resolution_override: int = None, + ) -> dict: + """ + Get proj conditioning for shape/texture stages (sparse-token aligned). + + Args: + image_cond_model: The proj image cond model for this stage. + image: List of PIL images. + coords: Sparse structure coordinates [N, 4] (batch_idx, x, y, z). + camera_angle_x: Camera horizontal FOV in radians. + distance: Camera distance. + mesh_scale: Mesh scale. + grid_resolution_override: Override the grid resolution if not None. + + Returns: + dict with 'cond' and 'neg_cond', each containing {'global': ..., 'proj': SparseTensor} + """ + print('Getting Projected Image Cond ...') + device = self.device + if self.low_vram: + image_cond_model.to(device) + + orig_grid_res = image_cond_model.grid_resolution + if grid_resolution_override is not None and grid_resolution_override != orig_grid_res: + image_cond_model.grid_resolution = grid_resolution_override + image_cond_model.proj_grid = image_cond_model.proj_grid.__class__( + grid_resolution=grid_resolution_override, + image_resolution=image_cond_model.proj_grid.image_resolution, + ).to(device) + + B = 1 + cam_angle = torch.tensor([camera_angle_x], device=device) + dist_tensor = torch.tensor([distance], device=device) + scale_tensor = torch.tensor([mesh_scale], device=device) + z_global, z_proj = image_cond_model( + image, camera_angle_x=cam_angle, distance=dist_tensor, mesh_scale=scale_tensor, + ) + grid_res = image_cond_model.grid_resolution + z_proj_grid = z_proj.reshape(B, grid_res, grid_res, grid_res, -1) + batch_indices = coords[:, 0].long() + x_coords = coords[:, 1].long() + y_coords = coords[:, 2].long() + z_coords = coords[:, 3].long() + z_proj_sparse = z_proj_grid[batch_indices, x_coords, y_coords, z_coords] + z_proj_st = SparseTensor(feats=z_proj_sparse, coords=coords) + + if grid_resolution_override is not None and grid_resolution_override != orig_grid_res: + image_cond_model.grid_resolution = orig_grid_res + image_cond_model.proj_grid = image_cond_model.proj_grid.__class__( + grid_resolution=orig_grid_res, + image_resolution=image_cond_model.proj_grid.image_resolution, + ).to(device) + + if self.low_vram: + image_cond_model.cpu() + return { + 'cond': {'global': z_global, 'proj': z_proj_st}, + 'neg_cond': {'global': torch.zeros_like(z_global), 'proj': SparseTensor(feats=torch.zeros_like(z_proj_sparse), coords=coords)}, + } + def preprocess_image(self, input: Image.Image) -> Image.Image: """ Preprocess the input image. @@ -1151,7 +1269,7 @@ class Trellis2ImageTo3DPipeline(Pipeline): # Get Image Cond self.load_image_cond_model() - # Multi-view conditioning happens inside get_cond() + # Multi-view conditioning happens inside get_cond() cond_512 = self.get_cond(images, 512, max_views = max_views) cond_1024 = self.get_cond(images, 1024, max_views = max_views) if pipeline_type != '512' else None @@ -3375,4 +3493,168 @@ class Trellis2ImageTo3DPipeline(Pipeline): else: return out_mesh, (shape_slat, None, res) else: - return out_mesh \ No newline at end of file + return out_mesh + + @torch.no_grad() + def run_pixal3d( + self, + image: Image.Image, + camera_params: dict, + num_samples: int = 1, + seed: int = 42, + sparse_structure_sampler_params: dict = {}, + shape_slat_sampler_params: dict = {}, + tex_slat_sampler_params: dict = {}, + preprocess_image: bool = True, + return_latent: bool = False, + pipeline_type: Optional[str] = None, + max_num_tokens: int = 49152, + generate_texture_slat: bool = False + ) -> List[MeshWithVoxel]: + """ + Run the Pixal3D pipeline (proj mode, cascade). + + Args: + image (Image.Image): The image prompt. + camera_params (dict): Camera parameters with keys: + - camera_angle_x (float): Horizontal FOV in radians. + - distance (float): Camera distance. + - mesh_scale (float): Mesh scale factor. + num_samples (int): The number of samples to generate. + seed (int): The random seed. + sparse_structure_sampler_params (dict): Additional parameters for the sparse structure sampler. + shape_slat_sampler_params (dict): Additional parameters for the shape SLat sampler. + tex_slat_sampler_params (dict): Additional parameters for the texture SLat sampler. + preprocess_image (bool): Whether to preprocess the image. + return_latent (bool): Whether to return the latent codes. + pipeline_type (str): The type of the pipeline. Options: '1024_cascade', '1536_cascade'. + max_num_tokens (int): The maximum number of tokens to use. + """ + # Check pipeline type + pipeline_type = pipeline_type or self.default_pipeline_type + + # Extract camera params + camera_angle_x = camera_params['camera_angle_x'] + distance = camera_params['distance'] + mesh_scale = camera_params.get('mesh_scale', 1.0) + + if preprocess_image: + image = self.preprocess_image(image) + torch.manual_seed(seed) + + # ---- Stage 1: Sparse Structure (proj) ---- + cond_ss = self.get_proj_cond_ss( + [image], + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + ) + ss_res = 32 + coords = self.sample_sparse_structure( + cond_ss, ss_res, + num_samples, sparse_structure_sampler_params + ) + del cond_ss + torch.cuda.empty_cache() + + # ---- Stage 2: Shape LR 512 (proj) ---- + cond_shape_lr = self.get_proj_cond_shape( + self.image_cond_model_shape_512, [image], coords, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + ) + lr_slat = self.sample_shape_slat( + cond_shape_lr, self.models['shape_slat_flow_model_512'], + coords, shape_slat_sampler_params + ) + del cond_shape_lr + torch.cuda.empty_cache() + + # ---- Stage 3a: Upsample LR → HR ---- + if self.low_vram: + self.models['shape_slat_decoder'].to(self.device) + self.models['shape_slat_decoder'].low_vram = True + hr_coords = self.models['shape_slat_decoder'].upsample(lr_slat, upsample_times=4) + if self.low_vram: + self.models['shape_slat_decoder'].cpu() + self.models['shape_slat_decoder'].low_vram = False + + lr_resolution = 512 + actual_hr_resolution = hr_resolution + while True: + grid_res = actual_hr_resolution // 16 + quant_coords = torch.cat([ + hr_coords[:, :1], + ((hr_coords[:, 1:] + 0.5) / lr_resolution * (grid_res - 1)).round().int(), + ], dim=1) + hr_coords_unique = quant_coords.unique(dim=0) + num_tokens = hr_coords_unique.shape[0] + if num_tokens < max_num_tokens or actual_hr_resolution == 1024: + break + actual_hr_resolution -= 128 + + actual_grid_res = actual_hr_resolution // 16 + del lr_slat, hr_coords, quant_coords + torch.cuda.empty_cache() + + # ---- Stage 3b: Shape HR (proj) ---- + cond_shape_hr = self.get_proj_cond_shape( + self.image_cond_model_shape_1024, [image], hr_coords_unique, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + grid_resolution_override=actual_grid_res, + ) + noise_hr = SparseTensor( + feats=torch.randn(hr_coords_unique.shape[0], self.models['shape_slat_flow_model_1024'].in_channels).to(self.device), + coords=hr_coords_unique, + ) + sampler_params_hr = {**self.shape_slat_sampler_params, **shape_slat_sampler_params} + flow_model_hr = self.models['shape_slat_flow_model_1024'] + if self.low_vram: + flow_model_hr.to(self.device) + hr_slat = self.shape_slat_sampler.sample( + flow_model_hr, + noise_hr, + **cond_shape_hr, + **sampler_params_hr, + verbose=True, + tqdm_desc=f"Sampling HR shape SLat (proj, {actual_hr_resolution})", + ).samples + if self.low_vram: + flow_model_hr.cpu() + std = torch.tensor(self.shape_slat_normalization['std'])[None].to(hr_slat.device) + mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(hr_slat.device) + shape_slat = hr_slat * std + mean + del cond_shape_hr, noise_hr, hr_slat, hr_coords_unique + torch.cuda.empty_cache() + + if generate_texture_slat: + # ---- Stage 4: Texture (proj) ---- + tex_grid_res = actual_hr_resolution // 16 + cond_tex = self.get_proj_cond_shape( + self.image_cond_model_tex_1024, [image], shape_slat.coords, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + grid_resolution_override=tex_grid_res, + ) + tex_slat = self.sample_tex_slat( + cond_tex, self.models['tex_slat_flow_model_1024'], + shape_slat, tex_slat_sampler_params + ) + del cond_tex + torch.cuda.empty_cache() + + # ---- Stage 5: Decode ---- + res = actual_hr_resolution + + if generate_texture_slat: + out_mesh = self.decode_latent(shape_slat, tex_slat, res) + else: + out_mesh = self.decode_latent(shape_slat, None, res) + if return_latent: + return out_mesh, (shape_slat, tex_slat, res) + else: + return out_mesh \ No newline at end of file diff --git a/trellis2/trainers/flow_matching/mixins/image_conditioned_proj.py b/trellis2/trainers/flow_matching/mixins/image_conditioned_proj.py new file mode 100644 index 0000000..b3198c2 --- /dev/null +++ b/trellis2/trainers/flow_matching/mixins/image_conditioned_proj.py @@ -0,0 +1,1530 @@ +""" +View-Aligned (Projection) Image Conditioned Mixin for TRELLIS2 + +This module implements DINOv3-based feature extraction with view-aligned projection, +supporting camera-aware 3D-to-2D feature mapping. +""" + +from typing import * +import os +import torch +import torch.nn as nn +import torch.nn.functional as F +from torchvision import transforms +from transformers import DINOv3ViTModel +import numpy as np +from PIL import Image, ImageDraw + +import torch.distributed as dist +from ....utils import dist_utils +from ....utils.dist_utils import read_file_dist + + +# ============================================================================= +# Projection Utilities +# ============================================================================= + +def project_points_to_image_batch( + points_3d: torch.Tensor, + transform_matrix: torch.Tensor, + camera_angle_x: torch.Tensor, + resolution: int = 518 +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Project 3D points to 2D image coordinates (batch processing). + + Args: + points_3d: torch.Tensor, shape [N, 3] or [B, N, 3], 3D point coordinates (in [-1, 1] range) + transform_matrix: torch.Tensor, shape [B, 4, 4], camera transformation matrix + camera_angle_x: torch.Tensor, shape [B], horizontal field of view angle (radians) + resolution: int, image resolution, default 518 + + Returns: + points_2d: torch.Tensor, shape [B, N, 2], image coordinates [x, y] + depth: torch.Tensor, shape [B, N], depth values + valid_mask: torch.Tensor, shape [B, N], mask for points within view + """ + device = points_3d.device + B = transform_matrix.shape[0] + + # Ensure inputs are torch.Tensor on correct device + if not isinstance(transform_matrix, torch.Tensor): + transform_matrix = torch.tensor(transform_matrix, dtype=torch.float32, device=device) + if not isinstance(points_3d, torch.Tensor): + points_3d = torch.tensor(points_3d, dtype=torch.float32, device=device) + if not isinstance(camera_angle_x, torch.Tensor): + camera_angle_x = torch.tensor(camera_angle_x, dtype=torch.float32, device=device) + + # Expand points_3d to batch dimension: [N, 3] -> [B, N, 3] + if points_3d.dim() == 2: + points_3d_batch = points_3d.unsqueeze(0).expand(B, -1, -1) + else: + points_3d_batch = points_3d + N = points_3d_batch.shape[1] + + # Add homogeneous coordinates: [B, N, 3] -> [B, N, 4] + ones = torch.ones(B, N, 1, device=device, dtype=points_3d_batch.dtype) + points_homogeneous = torch.cat([points_3d_batch, ones], dim=-1) # [B, N, 4] + + # Compute world to camera transformation matrix + world_to_camera = torch.linalg.inv(transform_matrix) # [B, 4, 4] + + # Batch transform to camera coordinate system: [B, N, 4] @ [B, 4, 4]^T -> [B, N, 3] + points_camera = torch.bmm(points_homogeneous, world_to_camera.transpose(-2, -1))[..., :3] # [B, N, 3] + + # Extract camera coordinates + x_cam = points_camera[..., 0] # [B, N] + y_cam = points_camera[..., 1] # [B, N] + z_cam = points_camera[..., 2] # [B, N] + + # Depth value (Z value in camera coordinate system, note Blender camera faces -Z direction) + depth = -z_cam # [B, N] + + # Compute camera intrinsics (batch processing) + sensor_width = 32.0 # mm + focal_length = 16.0 / torch.tan(camera_angle_x / 2.0) # [B] + focal_length_pixels = focal_length * resolution / sensor_width # [B] + + # Expand focal_length_pixels dimension for broadcasting: [B] -> [B, 1] + focal_length_pixels = focal_length_pixels.unsqueeze(1) # [B, 1] + + # Perspective projection to NDC coordinates + x_ndc = focal_length_pixels * x_cam / (-z_cam + 1e-8) # [B, N] + y_ndc = focal_length_pixels * y_cam / (-z_cam + 1e-8) # [B, N] + + # Convert to image coordinates (pixel coordinates) + x_pixel = x_ndc + resolution / 2.0 # [B, N] + y_pixel = -y_ndc + resolution / 2.0 # [B, N], flip Y axis + + # Create validity mask (points within image range and in front of camera) + valid_mask = ( + (x_pixel >= 0) & (x_pixel < resolution) & + (y_pixel >= 0) & (y_pixel < resolution) & + (depth > 0) # In front of camera + ) # [B, N] + + points_2d = torch.stack([x_pixel, y_pixel], dim=-1) # [B, N, 2] + + return points_2d, depth, valid_mask + + +def sample_features(fmap: torch.Tensor, queries_ndc: torch.Tensor) -> torch.Tensor: + """ + Sample features from feature map at specified NDC coordinates. + + Args: + fmap: torch.Tensor, shape [B, C, H, W], feature map + queries_ndc: torch.Tensor, shape [B, K, 2], normalized device coordinates + + Returns: + torch.Tensor, shape [B, C, K], sampled features + """ + B, C, H, W = fmap.shape + Bq, K, _ = queries_ndc.shape + assert Bq == B, "Batch size mismatch" + + # grid_sample requires (B, out_h, out_w, 2), here we want K points -> out_h=K, out_w=1 + grid = queries_ndc.view(B, K, 1, 2) # (B, K, 1, 2) + + # Bilinear interpolation, align_corners=False (consistent with [-1,1] pixel center convention) + feat = F.grid_sample( + fmap, grid, mode='bilinear', + align_corners=False, padding_mode='border' # border avoids out-of-bound becoming 0 + ) # (B, C, K, 1) + + return feat.squeeze(-1) # (B, C, K) + + +# ============================================================================= +# Projection Grid Module +# ============================================================================= + +class ProjGrid(nn.Module): + """ + 3D Grid Projection Module. + + Projects a 3D grid of points to 2D image coordinates and samples features + from the image feature map at those locations. + + This is the core module for view-aligned feature extraction. + """ + def __init__(self, grid_resolution: int = 16, image_resolution: int = 518): + super().__init__() + self.grid_resolution = grid_resolution + self.image_resolution = image_resolution + + # Create 3D grid points + one_dim = torch.linspace(-1, 1, grid_resolution) + x, y, z = torch.meshgrid(one_dim, one_dim, one_dim, indexing='ij') + grid_points = torch.stack((x, y, z), dim=-1) + + # Rotation matrix to align with Blender coordinate system + rotation_matrix = torch.tensor([ + [1.0, 0.0, 0.0], + [0.0, 0.0, -1.0], + [0.0, 1.0, 0.0] + ]) + grid_points = torch.matmul(grid_points, rotation_matrix.T) + grid_points = grid_points.reshape(-1, 3) + self.register_buffer('grid_points', grid_points) # [R³, 3] + + # Default front view transformation matrix + front_view_transform_matrix = torch.tensor([ + [1.0, 0.0, 0.0, 0.0], + [0.0, 0.0, -1.0, -2.0], + [0.0, 1.0, 0.0, 0.0], + [0.0, 0.0, 0.0, 1.0] + ]) + self.register_buffer("front_view_transform_matrix", front_view_transform_matrix) + + def forward( + self, + features_map: torch.Tensor, + camera_angle_x: torch.Tensor, + distance: torch.Tensor, + mesh_scale: torch.Tensor, + transform_matrix: Optional[torch.Tensor] = None, + BHWC: bool = True + ) -> torch.Tensor: + """ + Project 3D grid points to image and sample features. + + Args: + features_map: Feature map, shape [B, H, W, C] if BHWC else [B, C, H, W] + camera_angle_x: Camera FOV angle, shape [B] + distance: Camera distance, shape [B] + mesh_scale: Mesh scale factor, shape [B] + transform_matrix: Optional camera transform matrix, shape [B, 4, 4] + BHWC: Whether features_map is in BHWC format + + Returns: + Projected features, shape [B, grid_resolution³, C] + """ + if BHWC: + B, H, W, C = features_map.shape + else: + B, C, H, W = features_map.shape + + grid_points = self.grid_points + grid_points = grid_points.expand(B, -1, -1) + grid_points = grid_points / mesh_scale.unsqueeze(-1).unsqueeze(-1) / 2 # Scale alignment + assert transform_matrix is None, "transform_matrix is not None" + if transform_matrix is None: + transform_matrix = self.front_view_transform_matrix + transform_matrix = transform_matrix.expand(B, -1, -1).clone() + transform_matrix[:, 1, 3] = -distance # Set camera distance + + # Project to image coordinates (simulate Blender projection) + image_points, depth, valid_mask = project_points_to_image_batch( + grid_points, transform_matrix, camera_angle_x, self.image_resolution + ) + + # Normalize to [-1, 1] for grid_sample + image_points_norm = (image_points + 0.5) / self.image_resolution * 2 - 1 + + if BHWC: + features_map = features_map.permute(0, 3, 1, 2) # [B, C, H, W] + + # Sample features from DINOv3 patch feature map + x = sample_features(features_map, image_points_norm) # [B, C, K] + x = x.permute(0, 2, 1) # [B, K, C] + + return x + + def visualize_projection( + self, + image: torch.Tensor, + camera_angle_x: torch.Tensor, + distance: torch.Tensor, + mesh_scale: torch.Tensor, + transform_matrix: Optional[torch.Tensor] = None, + save_dir: Optional[str] = None, + prefix: str = "proj_vis", + ) -> List[Image.Image]: + """ + Visualize the projected 3D grid points on the input image. + + Args: + image: Input image tensor [B, C, H, W], assumed to be in [0, 1] range + camera_angle_x: Camera FOV angle, shape [B] + distance: Camera distance, shape [B] + mesh_scale: Mesh scale factor, shape [B] + transform_matrix: Optional camera transform matrix, shape [B, 4, 4] + save_dir: Directory to save visualizations (optional) + prefix: Prefix for saved files + + Returns: + List of PIL Images with projected points overlaid + """ + B = image.shape[0] + + # Get projected points + grid_points = self.grid_points.expand(B, -1, -1) + grid_points = grid_points / mesh_scale.unsqueeze(-1).unsqueeze(-1) / 2 + assert transform_matrix is None, "transform_matrix is not None" + if transform_matrix is None: + transform_matrix = self.front_view_transform_matrix + transform_matrix = transform_matrix.expand(B, -1, -1).clone() + transform_matrix[:, 1, 3] = -distance + + image_points, depth, valid_mask = project_points_to_image_batch( + grid_points, transform_matrix, camera_angle_x, self.image_resolution + ) + + # Convert image to PIL for visualization + vis_images = [] + for b in range(B): + # Convert tensor to PIL image + img_np = image[b].cpu().permute(1, 2, 0).numpy() + img_np = (img_np * 255).clip(0, 255).astype(np.uint8) + + # Resize to image_resolution if needed + pil_img = Image.fromarray(img_np) + if pil_img.size != (self.image_resolution, self.image_resolution): + pil_img = pil_img.resize((self.image_resolution, self.image_resolution), Image.LANCZOS) + + # Create a copy for drawing + vis_img = pil_img.copy() + draw = ImageDraw.Draw(vis_img) + + # Get points for this batch + pts = image_points[b].cpu().numpy() # [K, 2] + depths = depth[b].cpu().numpy() # [K] + mask = valid_mask[b].cpu().numpy() # [K] + + # Normalize depth for coloring + valid_depths = depths[mask] + if len(valid_depths) > 0: + d_min, d_max = valid_depths.min(), valid_depths.max() + if d_max - d_min > 1e-6: + depths_norm = (depths - d_min) / (d_max - d_min) + else: + depths_norm = np.ones_like(depths) * 0.5 + else: + depths_norm = np.ones_like(depths) * 0.5 + + # Draw projected points + R = self.grid_resolution + for i, (pt, d, m, dn) in enumerate(zip(pts, depths, mask, depths_norm)): + if not m: + continue + + x, y = pt + + # Color by depth (blue=near, red=far) + r = int(255 * dn) + g = int(255 * (1 - abs(2 * dn - 1))) + b_color = int(255 * (1 - dn)) + color = (r, g, b_color) + + # Draw small circle + radius = 2 + draw.ellipse( + [x - radius, y - radius, x + radius, y + radius], + fill=color, + outline=color + ) + + vis_images.append(vis_img) + + # Save if directory is specified + if save_dir is not None: + os.makedirs(save_dir, exist_ok=True) + save_path = os.path.join(save_dir, f"{prefix}_batch{b}.png") + vis_img.save(save_path) + print(f"Saved projection visualization to: {save_path}") + + return vis_images + + +# ============================================================================= +# DINOv3 Feature Extractor with Projection +# ============================================================================= + +class DinoV3ProjFeatureExtractor(nn.Module): + """ + DINOv3 Feature Extractor with View-Aligned Projection. + + This extractor produces both: + 1. Global features (CLS token + register tokens) in embed_dim + 2. View-aligned projected features (3D grid projected to 2D and sampled) + - Without NAF: [B, R³, embed_dim] + - With NAF: [B, R³, embed_dim * 2] (concat of lr and hr features) + + NOTE: proj_linear has been moved to per-block ProjectAttention / SparseProjectAttention. + This module now outputs raw DINOv3 features for proj (optionally concatenated with NAF-upsampled features). + + Args: + model_name: Name of the pretrained DINOv3 model + image_size: Input image size (default: 512) + grid_resolution: Resolution of the 3D projection grid (default: 16) + use_naf_upsample: Whether to use NAF to upsample features (default: False) + naf_target_size: Target spatial size for NAF upsampling (default: [128, 128]) + """ + def __init__( + self, + model_name: str, + image_size: int = 512, + grid_resolution: int = 16, + use_naf_upsample: bool = False, + naf_target_size: Optional[List[int]] = None + ): + super().__init__() + self.model_name = model_name + self.image_size = image_size + self.grid_resolution = grid_resolution + self.use_naf_upsample = use_naf_upsample + if naf_target_size is None: + self.naf_target_size = (128, 128) + elif isinstance(naf_target_size, int): + self.naf_target_size = (naf_target_size, naf_target_size) + else: + self.naf_target_size = tuple(naf_target_size) + + # Load DINOv3 model (frozen, no trainable params in this module) + self.model = DINOv3ViTModel.from_pretrained(model_name) + self.model.eval() + self.model.requires_grad_(False) + + # Image transform (only normalize, no resize - assume already resized) + self.transform = transforms.Compose([ + transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), + ]) + + # Get patch info + self.patch_size = self.model.config.patch_size + self.patch_number = image_size // self.patch_size + self.embed_dim = self.model.config.hidden_size + + # Projection grid for view-aligned features + self.proj_grid = ProjGrid( + grid_resolution=grid_resolution, + image_resolution=image_size + ) + + # NAF upsampler (frozen, no trainable params) + self.naf_model = None # Lazy-loaded on first use to avoid import if not needed + + # proj_channels: the output dimension of proj features + # Without NAF: embed_dim (e.g. 1024) + # With NAF: embed_dim * 2 (e.g. 2048, concat of lr and hr) + self.proj_channels = self.embed_dim * 2 if use_naf_upsample else self.embed_dim + + # NOTE: proj_linear removed — now lives in each denoiser block's ProjectAttention + + def _load_naf(self): + """Lazy-load pretrained NAF model.""" + if self.naf_model is None: + import torch.hub + device = next(self.model.parameters()).device + self.naf_model = torch.hub.load( + "valeoai/NAF", "naf", pretrained=True, device=device, trust_repo=True + ) + self.naf_model.eval() + self.naf_model.requires_grad_(False) + + def to(self, device): + super().to(device) + self.model.to(device) + self.proj_grid.to(device) + if self.naf_model is not None: + self.naf_model.to(device) + return self + + def cuda(self): + super().cuda() + self.model.cuda() + self.proj_grid.cuda() + if self.naf_model is not None: + self.naf_model.cuda() + return self + + def cpu(self): + super().cpu() + self.model.cpu() + self.proj_grid.cpu() + if self.naf_model is not None: + self.naf_model.cpu() + return self + + def extract_features(self, image: torch.Tensor) -> torch.Tensor: + """Extract features using DINOv3.""" + image = image.to(self.model.embeddings.patch_embeddings.weight.dtype) + hidden_states = self.model.embeddings(image, bool_masked_pos=None) + position_embeddings = self.model.rope_embeddings(image) + + for layer_module in self.model.layer: + hidden_states = layer_module( + hidden_states, + position_embeddings=position_embeddings, + ) + + return F.layer_norm(hidden_states, hidden_states.shape[-1:]) + + def forward( + self, + image: Union[torch.Tensor, List[Image.Image]], + camera_angle_x: Optional[torch.Tensor] = None, + distance: Optional[torch.Tensor] = None, + mesh_scale: Optional[torch.Tensor] = None, + transform_matrix: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Extract view-aligned features from the image. + + Args: + image: Input image tensor [B, C, H, W] or list of PIL images + camera_angle_x: Camera FOV angle in radians [B] + distance: Camera distance [B] + mesh_scale: Mesh scale factor [B] + transform_matrix: Optional camera transform matrix [B, 4, 4] + + Returns: + Tuple of (global_features, proj_features): + - global_features: [B, num_global_tokens, embed_dim] + - proj_features: [B, grid_resolution³, proj_channels] + where proj_channels = embed_dim (no NAF) or embed_dim*2 (with NAF) + """ + # Handle input types + if isinstance(image, torch.Tensor): + assert image.ndim == 4, "Image tensor should be batched (B, C, H, W)" + elif isinstance(image, list): + assert all(isinstance(i, Image.Image) for i in image), "Image list should be list of PIL images" + image = [i.resize((self.image_size, self.image_size), Image.LANCZOS) for i in image] + image = [np.array(i.convert('RGB')).astype(np.float32) / 255 for i in image] + image = [torch.from_numpy(i).permute(2, 0, 1).float() for i in image] + image = torch.stack(image).cuda() + else: + raise ValueError(f"Unsupported type of image: {type(image)}") + + B = image.shape[0] + + # Keep a copy of the unnormalized image for NAF guide + if self.use_naf_upsample: + image_for_naf = image.clone() # [B, 3, H, W], in [0, 1] range + + # Apply transform (ImageNet normalization) + image = self.transform(image) + + # Extract DINOv3 features (frozen, no gradients) + with torch.no_grad(): + z = self.extract_features(image) + + # Split into CLS token, register tokens, and patch tokens + z_clstoken = z[:, 0:1] # [B, 1, D] + num_reg = getattr(self.model.config, 'num_register_tokens', 4) + z_regtokens = z[:, 1:1+num_reg] # [B, num_reg, D] + z_patchtokens = z[:, 1+num_reg:] # [B, num_patches, D] + + # Reshape patch tokens to spatial grid: [B, h, w, D] + z_patchtokens_spatial = z_patchtokens.reshape( + B, self.patch_number, self.patch_number, -1 + ) # [B, h, w, D] + + if camera_angle_x is None or distance is None or mesh_scale is None: + raise ValueError("camera_angle_x, distance, and mesh_scale must be provided") + + # --- Low-resolution branch: sample from DINOv3 patch feature map --- + z_proj_lr = self.proj_grid( + z_patchtokens_spatial, + camera_angle_x, + distance, + mesh_scale, + transform_matrix + ) # [B, grid_res³, D] + + # --- High-resolution branch (NAF): upsample then sample --- + if self.use_naf_upsample: + self._load_naf() + # NAF expects: guide [B, 3, H, W], lr_features [B, C, h, w], target_size (H', W') + lr_features_bchw = z_patchtokens_spatial.permute(0, 3, 1, 2) # [B, D, h, w] + hr_features = self.naf_model( + image_for_naf, lr_features_bchw, self.naf_target_size + ) # [B, D, H', W'] + + # Sample from high-res feature map using same projection coordinates + z_proj_hr = self.proj_grid( + hr_features, + camera_angle_x, + distance, + mesh_scale, + transform_matrix, + BHWC=False # hr_features is [B, C, H', W'] + ) # [B, grid_res³, D] + + # Concatenate lr and hr: [B, grid_res³, D*2] + z_proj = torch.cat([z_proj_lr, z_proj_hr], dim=-1) + else: + z_proj = z_proj_lr # [B, grid_res³, D] + + # Combine global tokens + z_global = torch.cat([z_clstoken, z_regtokens], dim=1) # [B, 1+num_reg, D] + + # proj_linear has been moved to per-block ProjectAttention + # z_proj stays in proj_channels, each block will project independently + + return z_global, z_proj + + @torch.no_grad() + def visualize_projection( + self, + image: torch.Tensor, + camera_angle_x: torch.Tensor, + distance: torch.Tensor, + mesh_scale: torch.Tensor, + transform_matrix: Optional[torch.Tensor] = None, + save_dir: Optional[str] = None, + prefix: str = "proj_vis", + ) -> List[Image.Image]: + """ + Visualize the projected 3D grid points on the input image. + + This is a convenience method that delegates to ProjGrid.visualize_projection. + + Args: + image: Input image tensor [B, C, H, W], in [0, 1] range (before ImageNet normalization) + camera_angle_x: Camera FOV angle, shape [B] + distance: Camera distance, shape [B] + mesh_scale: Mesh scale factor, shape [B] + transform_matrix: Optional camera transform matrix, shape [B, 4, 4] + save_dir: Directory to save visualizations (optional) + prefix: Prefix for saved files + + Returns: + List of PIL Images with projected points overlaid + """ + return self.proj_grid.visualize_projection( + image=image, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + transform_matrix=transform_matrix, + save_dir=save_dir, + prefix=prefix, + ) + + +# ============================================================================= +# DINOv3 + VAE Gated Feature Extractor with Projection +# ============================================================================= + +class DinoV3VaeProjFeatureExtractor(nn.Module): + """ + DINOv3 + Flux VAE Feature Extractor with Gated Fusion and View-Aligned Projection. + + Produces three outputs for GatedProjectAttention: + 1. Global features (CLS + register tokens from DINOv3) for cross-attention + 2. Semantic proj features (DINOv3 patch tokens projected to 3D grid) + 3. Color proj features (Flux VAE latent projected to 3D grid) + + Both DINOv3 and VAE are frozen. The gated fusion happens inside each + denoiser block's GatedProjectAttention module (trainable gate + proj_linears). + + Args: + dino_model_name: Pretrained DINOv3 model name + vae_model_name: Pretrained Flux VAE model name + image_size: Input image size (default: 512) + grid_resolution: Resolution of the 3D projection grid (default: 16) + """ + def __init__( + self, + dino_model_name: str, + vae_model_name: str = "black-forest-labs/FLUX.1-dev", + image_size: int = 512, + grid_resolution: int = 16 + ): + super().__init__() + self.image_size = image_size + self.grid_resolution = grid_resolution + + # --- DINOv3 backbone (frozen) --- + self.dino_model = DINOv3ViTModel.from_pretrained(dino_model_name) + self.dino_model.eval() + self.dino_model.requires_grad_(False) + + self.dino_transform = transforms.Compose([ + transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), + ]) + + self.patch_size = self.dino_model.config.patch_size + self.patch_number = image_size // self.patch_size + self.embed_dim = self.dino_model.config.hidden_size # e.g. 1024 + + # --- Flux VAE encoder (frozen, lazy-loaded) --- + self.vae_model_name = vae_model_name + self._vae = None + self.vae_channels = 16 # Flux VAE outputs 16 channels + self.vae_downsample = 8 # Flux VAE downsamples by 8x + + # --- Projection grid (shared) --- + self.proj_grid = ProjGrid( + grid_resolution=grid_resolution, + image_resolution=image_size, + ) + + # Expose dimensions for denoiser block construction + self.dino_proj_channels = self.embed_dim # e.g. 1024 + self.vae_proj_channels = self.vae_channels # 16 + # proj_channels is kept for backward compat with _proj_channels in mixin + self.proj_channels = self.embed_dim + + def _load_vae(self): + """Lazy-load Flux VAE encoder.""" + if self._vae is not None: + return + from diffusers import AutoencoderKL + device = next(self.dino_model.parameters()).device + vae = AutoencoderKL.from_pretrained( + self.vae_model_name, + subfolder="vae", + torch_dtype=torch.float32, + ) + vae.eval() + vae.requires_grad_(False) + vae.to(device) + self._vae = vae + + def to(self, device): + super().to(device) + self.dino_model.to(device) + self.proj_grid.to(device) + if self._vae is not None: + self._vae.to(device) + return self + + def cuda(self): + super().cuda() + self.dino_model.cuda() + self.proj_grid.cuda() + if self._vae is not None: + self._vae.cuda() + return self + + def cpu(self): + super().cpu() + self.dino_model.cpu() + self.proj_grid.cpu() + if self._vae is not None: + self._vae.cpu() + return self + + def _extract_dino_features(self, image: torch.Tensor) -> torch.Tensor: + """Extract DINOv3 features from normalized image.""" + image = image.to(self.dino_model.embeddings.patch_embeddings.weight.dtype) + hidden_states = self.dino_model.embeddings(image, bool_masked_pos=None) + position_embeddings = self.dino_model.rope_embeddings(image) + for layer_module in self.dino_model.layer: + hidden_states = layer_module( + hidden_states, + position_embeddings=position_embeddings, + ) + return F.layer_norm(hidden_states, hidden_states.shape[-1:]) + + @torch.no_grad() + def _extract_vae_latent(self, image: torch.Tensor) -> torch.Tensor: + """Extract Flux VAE latent from unnormalized image [0,1].""" + self._load_vae() + image_normalized = image * 2.0 - 1.0 + image_normalized = image_normalized.to(self._vae.dtype) + posterior = self._vae.encode(image_normalized) + latent = posterior.latent_dist.mode() + latent = latent * self._vae.config.scaling_factor + return latent.float() # [B, 16, H/8, W/8] + + def forward( + self, + image: Union[torch.Tensor, List[Image.Image]], + camera_angle_x: Optional[torch.Tensor] = None, + distance: Optional[torch.Tensor] = None, + mesh_scale: Optional[torch.Tensor] = None, + transform_matrix: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Extract gated features from the image. + + Returns: + Tuple of (global_features, proj_semantic, proj_color): + - global_features: [B, num_global_tokens, embed_dim] (DINOv3 CLS + registers) + - proj_semantic: [B, grid_res³, embed_dim] (DINOv3 projected features) + - proj_color: [B, grid_res³, vae_channels] (VAE projected features) + """ + # Handle input types + if isinstance(image, torch.Tensor): + assert image.ndim == 4 + elif isinstance(image, list): + assert all(isinstance(i, Image.Image) for i in image) + image = [i.resize((self.image_size, self.image_size), Image.LANCZOS) for i in image] + image = [np.array(i.convert('RGB')).astype(np.float32) / 255 for i in image] + image = [torch.from_numpy(i).permute(2, 0, 1).float() for i in image] + image = torch.stack(image).cuda() + else: + raise ValueError(f"Unsupported type of image: {type(image)}") + + B = image.shape[0] + image_raw = image.clone() # Keep unnormalized copy for VAE + + if camera_angle_x is None or distance is None or mesh_scale is None: + raise ValueError("camera_angle_x, distance, and mesh_scale must be provided") + + with torch.no_grad(): + # --- DINOv3 branch --- + dino_input = self.dino_transform(image) + z = self._extract_dino_features(dino_input) + + z_clstoken = z[:, 0:1] + num_reg = getattr(self.dino_model.config, 'num_register_tokens', 4) + z_regtokens = z[:, 1:1+num_reg] + z_patchtokens = z[:, 1+num_reg:] + + z_patchtokens_spatial = z_patchtokens.reshape( + B, self.patch_number, self.patch_number, -1 + ) # [B, h, w, D] + + proj_semantic = self.proj_grid( + z_patchtokens_spatial, + camera_angle_x, distance, mesh_scale, transform_matrix, + ) # [B, R³, embed_dim] + + z_global = torch.cat([z_clstoken, z_regtokens], dim=1) # [B, 1+num_reg, D] + + # --- VAE branch --- + vae_latent = self._extract_vae_latent(image_raw) # [B, 16, H/8, W/8] + + proj_color = self.proj_grid( + vae_latent, + camera_angle_x, distance, mesh_scale, transform_matrix, + BHWC=False, # VAE latent is [B, C, H, W] + ) # [B, R³, 16] + + return z_global, proj_semantic, proj_color + + +# ============================================================================= +# Image Conditioned Mixin with Projection Support +# ============================================================================= + +class ImageConditionedProjMixin: + """ + Mixin for image-conditioned models with view-aligned projection. + + This mixin adds support for extracting view-aligned features from images + using camera parameters. + + Args: + image_cond_model: Configuration for the image conditioning model. + """ + def __init__(self, *args, image_cond_model: dict, **kwargs): + # Store config before super().__init__ which calls init_models_and_more + self.image_cond_model_config = image_cond_model + self.image_cond_model = None # Will be initialized in init_models_and_more + self.image_attn_mode = image_cond_model.get('image_attn_mode', + image_cond_model.get('args', {}).get('image_attn_mode', 'cross')) + super().__init__(*args, **kwargs) + + def _init_image_cond_model(self): + """Initialize the image conditioning model.""" + with dist_utils.local_master_first(): + model_name = self.image_cond_model_config['name'] + model_args = self.image_cond_model_config.get('args', {}) + + if model_name == 'DinoV3ProjFeatureExtractor': + self.image_cond_model = DinoV3ProjFeatureExtractor(**model_args) + elif model_name == 'DinoV3VaeProjFeatureExtractor': + self.image_cond_model = DinoV3VaeProjFeatureExtractor(**model_args) + else: + # Fallback to standard extractors + from . import image_conditioned + self.image_cond_model = getattr(image_conditioned, model_name)(**model_args) + + self.image_cond_model.cuda() + + # Expose proj_channels for denoiser to know the correct proj_in_channels + if hasattr(self.image_cond_model, 'proj_channels'): + self._proj_channels = self.image_cond_model.proj_channels + else: + self._proj_channels = getattr(self.image_cond_model, 'embed_dim', None) + # Expose vae_proj_channels for gated_proj mode + self._vae_proj_channels = getattr(self.image_cond_model, 'vae_proj_channels', None) + + def init_models_and_more(self, **kwargs): + """ + Override to handle image_cond_model initialization. + + Since proj_linear has been moved to per-block ProjectAttention in the denoiser, + image_cond_model no longer has any trainable parameters (DINOv3 backbone is frozen, + ProjGrid only has register_buffers). Therefore we do NOT add it to self.models + (which would trigger DDP wrapping and fail). We just initialize it and keep it + as a standalone module for inference. + """ + # Initialize image_cond_model first + if self.image_cond_model is None: + self._init_image_cond_model() + + # Keep a reference to the unwrapped module for attribute access + self._image_cond_module = self.image_cond_model # for .grid_resolution etc. + + # Log that image_cond has no trainable params + proj_params = [p for p in self.image_cond_model.parameters() if p.requires_grad] + if self.is_master: + if proj_params: + print(f'\nWARNING: image_cond_model has {len(proj_params)} trainable params, ' + f'but is NOT registered in self.models. These will NOT be trained!') + else: + print(f'\nimage_cond_model has no trainable parameters, skipping DDP/optimizer registration.') + + # Call base class to set up DDP, optimizer, EMA, etc. (without image_cond) + super().init_models_and_more(**kwargs) + + # ------------------------------------------------------------------ + # Checkpoint save/load overrides: skip DINOv3 backbone weights + # ------------------------------------------------------------------ + + # Keys in image_cond state_dict that belong to the frozen DINOv3 backbone. + # Everything under "model." is DINOv3; we only keep proj_grid.* + _IMAGE_COND_BACKBONE_PREFIX = 'model.' + + def _filter_image_cond_state_dict(self, state_dict: dict) -> dict: + """Keep only non-backbone keys (proj_grid, etc.) from image_cond state_dict.""" + return {k: v for k, v in state_dict.items() + if not k.startswith(self._IMAGE_COND_BACKBONE_PREFIX)} + + def _fill_denoiser_proj_linear_from_image_cond( + self, + denoiser_ckpt: dict, + denoiser_state_dict: dict, + image_cond_ckpt_path: Optional[str] = None, + ) -> dict: + """ + Fill missing per-block proj_linear weights in denoiser checkpoint + from the old-style image_cond proj_linear (broadcast to all blocks). + + Also handles shape mismatch when NAF is enabled: old proj_linear has shape + [model_ch, embed_dim] but new model expects [model_ch, embed_dim*2]. + In this case, the old weights are placed in the lr half and the hr half is zero-padded. + + Compatibility strategy: + 1. If denoiser_ckpt already contains per-block proj_linear keys with correct shape -> do nothing. + 2. If shape mismatch (embed_dim vs embed_dim*2) -> zero-pad the weight. + 3. If keys missing, try to load proj_linear from image_cond checkpoint -> broadcast (with optional pad). + + Args: + denoiser_ckpt: The loaded denoiser state dict + denoiser_state_dict: The model's current state dict (to find expected keys) + image_cond_ckpt_path: Path to image_cond checkpoint file (optional) + + Returns: + Updated denoiser_ckpt with proj_linear keys filled if needed + """ + if self.image_attn_mode != 'proj': + return denoiser_ckpt + + # Find all per-block proj_linear keys expected by the model + proj_linear_keys = [k for k in denoiser_state_dict.keys() + if '.cross_attn.proj_linear.' in k] + if not proj_linear_keys: + return denoiser_ckpt + + # --- Phase 1: Handle shape mismatch for existing keys (NAF upgrade) --- + for k in proj_linear_keys: + if k in denoiser_ckpt: + expected_shape = denoiser_state_dict[k].shape + actual_shape = denoiser_ckpt[k].shape + if expected_shape != actual_shape: + if k.endswith('.weight') and len(expected_shape) == 2: + # Weight shape: [out_features, in_features] + # Old: [model_ch, embed_dim], New: [model_ch, embed_dim*2] + out_f, new_in_f = expected_shape + _, old_in_f = actual_shape + if new_in_f > old_in_f and out_f == actual_shape[0]: + if self.is_master: + print(f'\n [NAF Compat] Padding proj_linear weight {k}: ' + f'{actual_shape} -> {expected_shape} (zero-pad hr half)') + new_w = torch.zeros(expected_shape, dtype=denoiser_ckpt[k].dtype, + device=denoiser_ckpt[k].device) + new_w[:, :old_in_f] = denoiser_ckpt[k] + denoiser_ckpt[k] = new_w + else: + if self.is_master: + print(f'\n Warning: proj_linear {k} shape mismatch ' + f'{actual_shape} vs {expected_shape}, using model init') + denoiser_ckpt[k] = denoiser_state_dict[k] + # bias shape should match (out_features only), no padding needed + + # --- Phase 2: Handle completely missing keys --- + missing_proj_keys = [k for k in proj_linear_keys if k not in denoiser_ckpt] + if not missing_proj_keys: + return denoiser_ckpt + + if self.is_master: + print(f'\n [Compat] Denoiser ckpt missing {len(missing_proj_keys)} per-block proj_linear keys.') + print(f' Attempting to load from image_cond proj_linear: {image_cond_ckpt_path}') + + # Try to find proj_linear weights from image_cond checkpoint + old_proj_linear_w = None + old_proj_linear_b = None + + if image_cond_ckpt_path is not None: + import os as _os + if _os.path.exists(image_cond_ckpt_path): + try: + ic_ckpt = torch.load(image_cond_ckpt_path, map_location=self.device, weights_only=True) + old_proj_linear_w = ic_ckpt.get('proj_linear.weight') + old_proj_linear_b = ic_ckpt.get('proj_linear.bias') + except Exception as e: + if self.is_master: + print(f' Warning: Failed to load image_cond ckpt: {e}') + + if old_proj_linear_w is None: + raise RuntimeError( + f'Denoiser checkpoint is missing per-block proj_linear keys ' + f'(e.g. {missing_proj_keys[0]}), and no image_cond proj_linear ' + f'was found to broadcast from. Cannot proceed.' + ) + + if self.is_master: + print(f' Found image_cond proj_linear: weight {old_proj_linear_w.shape}, bias {old_proj_linear_b.shape}') + print(f' Broadcasting to {len(missing_proj_keys)} per-block keys...') + + for k in missing_proj_keys: + if k.endswith('.weight'): + expected_shape = denoiser_state_dict[k].shape + if expected_shape != old_proj_linear_w.shape: + # Pad for NAF: [model_ch, embed_dim] -> [model_ch, embed_dim*2] + out_f, new_in_f = expected_shape + _, old_in_f = old_proj_linear_w.shape + if new_in_f > old_in_f and out_f == old_proj_linear_w.shape[0]: + new_w = torch.zeros(expected_shape, dtype=old_proj_linear_w.dtype) + new_w[:, :old_in_f] = old_proj_linear_w + denoiser_ckpt[k] = new_w + else: + denoiser_ckpt[k] = denoiser_state_dict[k] + else: + denoiser_ckpt[k] = old_proj_linear_w.clone() + elif k.endswith('.bias'): + denoiser_ckpt[k] = old_proj_linear_b.clone() + + return denoiser_ckpt + + def _master_params_to_state_dicts(self, master_params): + """Override to skip image_cond checkpoint entirely. + + image_cond model no longer has trainable parameters: + - proj_linear has been moved to per-block ProjectAttention in the denoiser + - DINOv3 backbone is frozen and loaded from pretrained weights + - ProjGrid only contains fixed register_buffers (grid_points, front_view_transform_matrix) + So there is nothing worth saving for image_cond. + """ + state_dicts = super()._master_params_to_state_dicts(master_params) + state_dicts.pop('image_cond', None) + return state_dicts + + def load(self, load_dir, step=0): + """ + Override to handle: + 1. Old checkpoints that don't have image_cond_step*.pt + 2. Partial image_cond checkpoints (only proj_linear + proj_grid, no DINOv3 backbone) + """ + import os as _os + + if self.is_master: + print(f'\nLoading checkpoint from step {step}...', end='') + + model_ckpts = {} + for name, model in self.models.items(): + ckpt_path = _os.path.join(load_dir, 'ckpts', f'{name}_step{step:07d}.pt') + + if name == 'image_cond': + # --- handle missing or partial image_cond checkpoint --- + if not _os.path.exists(ckpt_path): + if self.is_master: + print(f'\n image_cond checkpoint not found at {ckpt_path}, using freshly initialised weights.') + model_ckpts[name] = model.state_dict() + continue + + try: + model_ckpt = torch.load( + read_file_dist(ckpt_path), + map_location=self.device, weights_only=True) + except Exception as e: + if self.is_master: + print(f'\n Failed to load image_cond checkpoint: {e}. Using freshly initialised weights.') + model_ckpts[name] = model.state_dict() + continue + + # Partial ckpt (no backbone) → load with strict=False + missing, unexpected = model.load_state_dict(model_ckpt, strict=False) + # All missing keys should be the frozen DINOv3 backbone; verify + non_backbone_missing = [k for k in missing + if not k.startswith(self._IMAGE_COND_BACKBONE_PREFIX)] + if non_backbone_missing and self.is_master: + print(f'\n Warning: unexpected missing keys in image_cond ckpt: {non_backbone_missing}') + if unexpected and self.is_master: + print(f'\n Warning: unexpected keys in image_cond ckpt: {unexpected}') + + # Build a full state_dict for master_params sync + full_sd = model.state_dict() + full_sd.update(model_ckpt) + model_ckpts[name] = full_sd + else: + model_ckpt = torch.load( + read_file_dist(ckpt_path), + map_location=self.device, weights_only=True) + # For denoiser: handle old ckpts missing per-block proj_linear + if name == 'denoiser': + ic_ckpt_path = _os.path.join(load_dir, 'ckpts', f'image_cond_step{step:07d}.pt') + model_ckpt = self._fill_denoiser_proj_linear_from_image_cond( + model_ckpt, model.state_dict(), ic_ckpt_path) + model_ckpts[name] = model_ckpt + model.load_state_dict(model_ckpt) + + self._state_dicts_to_master_params(self.master_params, model_ckpts) + del model_ckpts + + if self.is_master: + for i, ema_rate in enumerate(self.ema_rate): + ema_ckpts = {} + for name, model in self.models.items(): + ema_path = _os.path.join( + load_dir, 'ckpts', + f'{name}_ema{ema_rate}_step{step:07d}.pt') + if name == 'image_cond': + if not _os.path.exists(ema_path): + ema_ckpts[name] = model.state_dict() + continue + try: + ema_ckpt = torch.load(ema_path, map_location=self.device, weights_only=True) + except Exception: + ema_ckpts[name] = model.state_dict() + continue + full_sd = model.state_dict() + full_sd.update(ema_ckpt) + ema_ckpts[name] = full_sd + else: + ema_ckpt = torch.load(ema_path, map_location=self.device, weights_only=True) + if name == 'denoiser': + ic_ema_path = _os.path.join( + load_dir, 'ckpts', + f'image_cond_ema{ema_rate}_step{step:07d}.pt') + ema_ckpt = self._fill_denoiser_proj_linear_from_image_cond( + ema_ckpt, model.state_dict(), ic_ema_path) + ema_ckpts[name] = ema_ckpt + self._state_dicts_to_master_params(self.ema_params[i], ema_ckpts) + del ema_ckpts + + misc_ckpt = torch.load( + read_file_dist(_os.path.join(load_dir, 'ckpts', f'misc_step{step:07d}.pt')), + map_location=torch.device('cpu'), weights_only=False) + # Optimizer state may mismatch when loading old checkpoints that were + # saved before image_cond was added to self.models, or when the number + # of trainable parameters changed (e.g. backbone freeze, NAF upgrade). + # In that case we skip restoring optimizer state and let it re-initialise. + try: + self.optimizer.load_state_dict(misc_ckpt['optimizer']) + # Verify optimizer state shapes match parameters. + # load_state_dict may succeed even when shapes mismatch (keys are + # integer indices), causing a crash later in optimizer.step(). + _shape_ok = True + for group in self.optimizer.param_groups: + for p in group['params']: + state = self.optimizer.state.get(p) + if state is not None: + for sv in state.values(): + if isinstance(sv, torch.Tensor) and sv.shape != () and sv.shape != p.shape: + _shape_ok = False + break + if not _shape_ok: + break + if not _shape_ok: + break + if not _shape_ok: + if self.is_master: + print(f'\n Warning: optimizer state shape mismatch (likely NAF upgrade). ' + f'Optimizer will start fresh.') + self.optimizer.state.clear() + except (ValueError, RuntimeError) as e: + if self.is_master: + print(f'\n Warning: could not load optimizer state ({e}). ' + f'Optimizer will start fresh.') + self.step = misc_ckpt['step'] + self.data_sampler.load_state_dict(misc_ckpt['data_sampler']) + if self.mix_precision_mode == 'amp' and self.mix_precision_dtype == torch.float16: + self.scaler.load_state_dict(misc_ckpt['scaler']) + elif self.mix_precision_mode == 'inflat_all' and self.mix_precision_dtype == torch.float16: + self.log_scale = misc_ckpt['log_scale'] + if self.lr_scheduler_config is not None: + self.lr_scheduler.load_state_dict(misc_ckpt['lr_scheduler']) + if self.elastic_controller_config is not None: + self.elastic_controller.load_state_dict(misc_ckpt['elastic_controller']) + if self.grad_clip is not None and not isinstance(self.grad_clip, float): + self.grad_clip.load_state_dict(misc_ckpt['grad_clip']) + del misc_ckpt + + if self.world_size > 1: + dist.barrier() + if self.is_master: + print(' Done.') + + if self.world_size > 1: + self.check_ddp() + + def finetune_from(self, finetune_ckpt): + """ + Override to tolerate DINOv3 backbone keys missing from image_cond checkpoint. + For image_cond, the checkpoint only stores proj_linear + proj_grid (no backbone), + so we treat all backbone keys as allowed-missing. + """ + ALLOWED_MISSING_KEYS = {'rope_phases'} + + if self.is_master: + print('\nFinetuning from:') + for name, path in finetune_ckpt.items(): + print(f' - {name}: {path}') + + model_ckpts = {} + for name, model in self.models.items(): + model_state_dict = model.state_dict() + if name in finetune_ckpt: + model_ckpt = torch.load( + read_file_dist(finetune_ckpt[name]), + map_location=self.device, weights_only=True) + + model_ckpt = self._remap_checkpoint_keys(model_ckpt, model_state_dict) + + for k, v in model_ckpt.items(): + if k not in model_state_dict: + if self.is_master: + print(f'Warning: {k} not found in model_state_dict, skipped.') + model_ckpt[k] = None + elif model_ckpt[k].shape != model_state_dict[k].shape: + # For proj_linear weights, try zero-pad instead of skipping + # This handles NAF upgrade: [model_ch, embed_dim] -> [model_ch, embed_dim*2] + if '.cross_attn.proj_linear.weight' in k and len(model_ckpt[k].shape) == 2: + old_shape = model_ckpt[k].shape + new_shape = model_state_dict[k].shape + if new_shape[0] == old_shape[0] and new_shape[1] > old_shape[1]: + if self.is_master: + print(f'Info: Zero-padding proj_linear weight {k}: {old_shape} -> {new_shape}') + new_w = torch.zeros(new_shape, dtype=model_ckpt[k].dtype) + new_w[:, :old_shape[1]] = model_ckpt[k] + model_ckpt[k] = new_w + else: + if self.is_master: + print(f'Warning: {k} shape mismatch, {old_shape} vs {new_shape}, skipped.') + model_ckpt[k] = model_state_dict[k] + else: + if self.is_master: + print(f'Warning: {k} shape mismatch, {model_ckpt[k].shape} vs {model_state_dict[k].shape}, skipped.') + model_ckpt[k] = model_state_dict[k] + model_ckpt = {k: v for k, v in model_ckpt.items() if v is not None} + + missing_keys = set(model_state_dict.keys()) - set(model_ckpt.keys()) + + # For denoiser: fill per-block proj_linear from image_cond if missing + if name == 'denoiser': + ic_path = finetune_ckpt.get('image_cond') + proj_linear_missing = {k for k in missing_keys if '.cross_attn.proj_linear.' in k} + if proj_linear_missing: + model_ckpt = self._fill_denoiser_proj_linear_from_image_cond( + model_ckpt, model_state_dict, ic_path) + # Recalculate missing_keys after filling + missing_keys = set(model_state_dict.keys()) - set(model_ckpt.keys()) + + # For image_cond, DINOv3 backbone keys are expected to be missing + allowed = set(ALLOWED_MISSING_KEYS) + if name == 'image_cond': + backbone_missing = {k for k in missing_keys + if k.startswith(self._IMAGE_COND_BACKBONE_PREFIX)} + allowed |= backbone_missing + if backbone_missing and self.is_master: + print(f'Info: image_cond: {len(backbone_missing)} DINOv3 backbone keys ' + f'not in ckpt (expected, using pretrained weights)') + # Old ckpts may have proj_linear.* which has moved to denoiser blocks + proj_linear_missing = {k for k in missing_keys if k.startswith('proj_linear.')} + allowed |= proj_linear_missing + + unexpected_missing = missing_keys - allowed + if unexpected_missing and self.is_master: + print(f'Error: Missing keys in checkpoint: {unexpected_missing}') + raise RuntimeError(f'Missing keys in checkpoint: {unexpected_missing}') + if missing_keys & ALLOWED_MISSING_KEYS and self.is_master: + print(f'Info: Using model initialized values for: {missing_keys & ALLOWED_MISSING_KEYS}') + + for k in missing_keys: + model_ckpt[k] = model_state_dict[k] + + model_ckpts[name] = model_ckpt + model.load_state_dict(model_ckpt) + else: + if self.is_master: + print(f'Warning: {name} not found in finetune_ckpt, skipped.') + model_ckpts[name] = model_state_dict + + self._state_dicts_to_master_params(self.master_params, model_ckpts) + if self.is_master: + for i, ema_rate in enumerate(self.ema_rate): + self._state_dicts_to_master_params(self.ema_params[i], model_ckpts) + del model_ckpts + + if self.world_size > 1: + dist.barrier() + if self.is_master: + print('Done.') + + if self.world_size > 1: + self.check_ddp() + + def encode_image_proj( + self, + image: torch.Tensor, + camera_angle_x: Optional[torch.Tensor] = None, + distance: Optional[torch.Tensor] = None, + mesh_scale: Optional[torch.Tensor] = None, + transform_matrix: Optional[torch.Tensor] = None, + coords: Optional[torch.Tensor] = None, + ) -> Tuple[Dict[str, torch.Tensor], Dict[str, torch.Tensor]]: + """ + Encode the image with view-aligned projection. + + Supports both 'proj' mode (DINOv3 only, 2 outputs) and + 'gated_proj' mode (DINOv3 + VAE, 3 outputs). + """ + if self.image_cond_model is None: + self._init_image_cond_model() + + outputs = self.image_cond_model( + image, + camera_angle_x=camera_angle_x, + distance=distance, + mesh_scale=mesh_scale, + transform_matrix=transform_matrix, + ) + + is_gated = self.image_attn_mode == 'gated_proj' + + if is_gated: + cond_global, cond_proj_semantic, cond_proj_color = outputs + else: + cond_global, cond_proj = outputs + + # If coords provided, extract features at sparse positions + if coords is not None: + B = cond_global.shape[0] + module = getattr(self, '_image_cond_module', self.image_cond_model) + grid_res = module.grid_resolution + batch_indices = coords[:, 0].long() + x_coords = coords[:, 1].long() + y_coords = coords[:, 2].long() + z_coords = coords[:, 3].long() + + if is_gated: + cond_proj_semantic = cond_proj_semantic.reshape(B, grid_res, grid_res, grid_res, -1) + cond_proj_semantic = cond_proj_semantic[batch_indices, x_coords, y_coords, z_coords] + cond_proj_color = cond_proj_color.reshape(B, grid_res, grid_res, grid_res, -1) + cond_proj_color = cond_proj_color[batch_indices, x_coords, y_coords, z_coords] + else: + cond_proj = cond_proj.reshape(B, grid_res, grid_res, grid_res, -1) + cond_proj = cond_proj[batch_indices, x_coords, y_coords, z_coords] + + if is_gated: + cond = { + 'global': cond_global, + 'proj_semantic': cond_proj_semantic, + 'proj_color': cond_proj_color, + } + uncond = { + 'global': torch.zeros_like(cond_global), + 'proj_semantic': torch.zeros_like(cond_proj_semantic), + 'proj_color': torch.zeros_like(cond_proj_color), + } + else: + cond = {'global': cond_global, 'proj': cond_proj} + uncond = {'global': torch.zeros_like(cond_global), 'proj': torch.zeros_like(cond_proj)} + + return cond, uncond + + @torch.no_grad() + def encode_image(self, image: Union[torch.Tensor, List[Image.Image]]) -> torch.Tensor: + """ + Encode the image (standard mode without projection). + """ + if self.image_cond_model is None: + self._init_image_cond_model() + + if self.image_attn_mode == 'proj': + # For proj mode, return dict + global_feat, proj_feat = self.image_cond_model(image) + return {'global': global_feat, 'proj': proj_feat} + else: + # Standard mode + features = self.image_cond_model(image) + return features + + def _extract_camera_info(self, kwargs): + """ + Extract camera info from kwargs. + + Supports two formats: + 1. 'camera_info' dict: {'camera_angle_x': ..., 'distance': ..., 'mesh_scale': ..., 'transform_matrix': ..., 'coords': ...} + 2. Flat fields: 'camera_angle_x', 'camera_distance', 'mesh_scale', 'transform_matrix', 'coords' in kwargs + + Returns: + camera_info dict or None if not available + """ + if 'camera_info' in kwargs: + return kwargs.pop('camera_info') + + # Try to extract from flat fields (as returned by ViewImageConditionedMixin) + camera_angle_x = kwargs.pop('camera_angle_x', None) + camera_distance = kwargs.pop('camera_distance', None) + mesh_scale = kwargs.pop('mesh_scale', None) + transform_matrix = kwargs.pop('transform_matrix', None) + coords = kwargs.pop('coords', None) + + if camera_angle_x is not None and camera_distance is not None and mesh_scale is not None: + return { + 'camera_angle_x': camera_angle_x, + 'distance': camera_distance, + 'mesh_scale': mesh_scale, + 'transform_matrix': transform_matrix, + 'coords': coords, + } + + return None + + def get_cond(self, cond, **kwargs): + """Get the conditioning data.""" + kwargs.pop('view_idx', None) + + if self.image_attn_mode in ('proj', 'gated_proj'): + # Handle projection mode (both standard proj and gated_proj) + camera_info = self._extract_camera_info(kwargs) + if camera_info is not None: + coords = camera_info.get('coords') + cond, neg_cond = self.encode_image_proj( + cond, + camera_angle_x=camera_info.get('camera_angle_x'), + distance=camera_info.get('distance'), + mesh_scale=camera_info.get('mesh_scale'), + transform_matrix=camera_info.get('transform_matrix'), + coords=coords, + ) + + # For sparse mode (coords provided), handle CFG dropout ourselves + if coords is not None and hasattr(self, 'p_uncond') and self.p_uncond > 0: + import numpy as np + B = cond['global'].shape[0] + mask = np.random.rand(B) < self.p_uncond + + global_tensor = cond['global'] + global_mask_shape = [B] + [1] * (global_tensor.ndim - 1) + global_mask = torch.tensor(mask, device=global_tensor.device).reshape(global_mask_shape) + cond['global'] = torch.where(global_mask, neg_cond['global'], cond['global']) + + batch_indices = coords[:, 0].long() + # Handle all sparse proj keys (proj, or proj_semantic + proj_color) + for key in list(cond.keys()): + if key.startswith('proj'): + device = cond[key].device + sparse_mask = torch.tensor(mask, device=device)[batch_indices].reshape(-1, 1) + cond[key] = torch.where(sparse_mask, neg_cond[key], cond[key]) + + return cond + else: + kwargs['neg_cond'] = neg_cond + else: + cond = self.encode_image(cond) + if isinstance(cond, dict) and 'global' in cond: + kwargs['neg_cond'] = {k: torch.zeros_like(v) for k, v in cond.items()} + else: + kwargs['neg_cond'] = torch.zeros_like(cond) + else: + cond = self.encode_image(cond) + kwargs['neg_cond'] = torch.zeros_like(cond) + + cond = super().get_cond(cond, **kwargs) + return cond + + def get_inference_cond(self, cond, **kwargs): + """Get the conditioning data for inference.""" + kwargs.pop('view_idx', None) + + if self.image_attn_mode in ('proj', 'gated_proj'): + camera_info = self._extract_camera_info(kwargs) + if camera_info is not None: + cond, neg_cond = self.encode_image_proj( + cond, + camera_angle_x=camera_info.get('camera_angle_x'), + distance=camera_info.get('distance'), + mesh_scale=camera_info.get('mesh_scale'), + transform_matrix=camera_info.get('transform_matrix'), + coords=camera_info.get('coords'), + ) + kwargs['neg_cond'] = neg_cond + else: + cond = self.encode_image(cond) + if isinstance(cond, dict) and 'global' in cond: + kwargs['neg_cond'] = {k: torch.zeros_like(v) for k, v in cond.items()} + else: + kwargs['neg_cond'] = torch.zeros_like(cond) + else: + cond = self.encode_image(cond) + kwargs['neg_cond'] = torch.zeros_like(cond) + + cond = super().get_inference_cond(cond, **kwargs) + return cond + + def vis_cond(self, cond, **kwargs): + """Visualize the conditioning data.""" + return {'image': {'value': cond, 'type': 'image'}} + + @torch.no_grad() + def visualize_projection_test( + self, + cond: torch.Tensor, + save_dir: str, + prefix: str = "proj_vis", + **kwargs + ) -> Optional[List[Image.Image]]: + """ + Visualize projection points on the condition images. + + This should be called once before training starts to verify the projection is correct. + + Args: + cond: Condition image tensor [B, C, H, W], in [0, 1] range + save_dir: Directory to save visualizations + prefix: Prefix for saved files + **kwargs: Should contain camera_angle_x, camera_distance, mesh_scale, transform_matrix + + Returns: + List of PIL Images with projected points overlaid, or None if not in proj mode + """ + if self.image_attn_mode != 'proj': + return None + + if self.image_cond_model is None: + self._init_image_cond_model() + + # Use _image_cond_module for attribute access (image_cond_model may be DDP-wrapped) + module = getattr(self, '_image_cond_module', self.image_cond_model) + + # Check if the model has visualization capability + if not hasattr(module, 'visualize_projection'): + print("Warning: image_cond_model does not support visualize_projection") + return None + + # Extract camera info + camera_info = self._extract_camera_info(kwargs) + if camera_info is None: + print("Warning: No camera info available for projection visualization") + return None + + return module.visualize_projection( + image=cond, + camera_angle_x=camera_info.get('camera_angle_x'), + distance=camera_info.get('distance'), + mesh_scale=camera_info.get('mesh_scale'), + transform_matrix=camera_info.get('transform_matrix'), + save_dir=save_dir, + prefix=prefix, + ) diff --git a/trellis2/utils/camera.py b/trellis2/utils/camera.py new file mode 100644 index 0000000..3a9a87c --- /dev/null +++ b/trellis2/utils/camera.py @@ -0,0 +1,40 @@ +import torch +import numpy as np +import math + +def compute_f_pixels(camera_angle_x: float, resolution: int) -> float: + focal_length = 16.0 / torch.tan(torch.tensor(camera_angle_x / 2.0)) + f_pixels = focal_length * resolution / 32.0 + return float(f_pixels.item()) + +def distance_from_fov(camera_angle_x, grid_point, target_point, mesh_scale, image_resolution): + rotation_matrix = torch.tensor([[1.0, 0.0, 0.0], [0.0, 0.0, -1.0], [0.0, 1.0, 0.0]]) + gp = grid_point.to(torch.float32) @ rotation_matrix.T + gp = gp / mesh_scale / 2 + xw, yw, zw = gp[0].item(), gp[1].item(), gp[2].item() + xt, yt = float(target_point[0].item()), float(target_point[1].item()) + f_pixels = compute_f_pixels(camera_angle_x, image_resolution) + x_ndc = xt - image_resolution / 2.0 + y_ndc = -(yt - image_resolution / 2.0) + distance_x = f_pixels * xw / x_ndc - yw + return {"distance_from_x": float(distance_x), "f_pixels": float(f_pixels)} + +def get_camera_params_wild_moge(pil_image, moge_model, device="cuda", mesh_scale=1.0, extend_pixel=0, image_resolution=512): + #pil_image = Image.open(image_path).convert("RGB") + width, height = pil_image.size + image_np = np.array(pil_image).astype(np.float32) / 255.0 + image_tensor = torch.from_numpy(image_np).permute(2, 0, 1).to(device) + with torch.no_grad(): + output = moge_model.infer(image_tensor) + intrinsics = output["intrinsics"].squeeze().cpu().numpy() + fx_normalized = intrinsics[0, 0] + fx = fx_normalized * width + camera_angle_x = 2 * math.atan(width / (2 * fx)) + + grid_point = torch.tensor([-1.0, 0.0, 0.0]) + distance = distance_from_fov( + camera_angle_x, grid_point, + torch.tensor([0 - extend_pixel, image_resolution - 1 + extend_pixel]), + mesh_scale, image_resolution + )["distance_from_x"] + return {'camera_angle_x': camera_angle_x, 'distance': distance, 'mesh_scale': mesh_scale} \ No newline at end of file diff --git a/wheels/Windows/Torch2100/CUDA 13.1/natten-0.21.6-cp313-cp313-win_amd64.whl b/wheels/Windows/Torch2100/CUDA 13.1/natten-0.21.6-cp313-cp313-win_amd64.whl new file mode 100644 index 0000000..9597649 Binary files /dev/null and b/wheels/Windows/Torch2100/CUDA 13.1/natten-0.21.6-cp313-cp313-win_amd64.whl differ diff --git a/wheels/Windows/Torch2100/CUDA 13.1/nvdiffrec_render-0.0.0-cp313-cp313-win_amd64.whl b/wheels/Windows/Torch2100/CUDA 13.1/nvdiffrec_render-0.0.0-cp313-cp313-win_amd64.whl new file mode 100644 index 0000000..9fd1dec Binary files /dev/null and b/wheels/Windows/Torch2100/CUDA 13.1/nvdiffrec_render-0.0.0-cp313-cp313-win_amd64.whl differ