XIAN-FEM-2026June/3D opticsfem-master/tools/validate_and_fix_pbcmesh.py

305 lines
9.5 KiB
Python

#!/usr/bin/env python3
"""Validate PBCmesh.dat and repair NormOfFace for domains 1-9 (PBC model)."""
from __future__ import annotations
import argparse
import math
import sys
from pathlib import Path
import numpy as np
# PBC single-period case: DomainOfTri 1..9
ALL_FACE_DOMAINS = list(range(1, 10))
SBC_OUT = {3, 9}
SBC_INC = {8}
def read_tag(lines: list[str], i: int, expected: str) -> int:
while i < len(lines) and not lines[i].strip():
i += 1
if i >= len(lines) or lines[i].strip() != expected:
got = lines[i].strip() if i < len(lines) else "EOF"
raise ValueError(f"Expected tag '{expected}' at line {i + 1}, got '{got}'")
return i + 1
def face_vertices(V: np.ndarray, T: np.ndarray, num_tet: int, num_face: int) -> np.ndarray:
"""Return 3x3 array of face vertex coordinates (reference tet local face)."""
x = V[T[num_tet] - 1, 0]
y = V[T[num_tet] - 1, 1]
z = V[T[num_tet] - 1, 2]
if num_face == 1:
idx = [0, 1, 2]
elif num_face == 2:
idx = [0, 1, 3]
elif num_face == 3:
idx = [0, 2, 3]
elif num_face == 4:
idx = [1, 2, 3]
else:
raise ValueError(f"bad face {num_face}")
pts = np.column_stack([x[idx], y[idx], z[idx]])
return pts
def geometric_face_normal(pts: np.ndarray) -> np.ndarray:
n = np.cross(pts[1] - pts[0], pts[2] - pts[0])
norm = np.linalg.norm(n)
if norm < 1e-30:
raise ValueError("degenerate face normal")
return n / norm
def matlab_out_normal(domain: int, tet_vertex_xyz: np.ndarray) -> np.ndarray:
"""assembly_out.m normal for one SBC triangle."""
if domain == 3:
return np.array([0.0, 0.0, -1.0])
x = tet_vertex_xyz[:, 0]
y = tet_vertex_xyz[:, 1]
n = np.array([x.sum() / 3.0, y.sum() / 3.0, 0.0])
norm = np.linalg.norm(n)
if norm < 1e-30:
raise ValueError(f"Degenerate assembly_out normal domain {domain}")
return n / norm
def matlab_inc_normal() -> np.ndarray:
return np.array([0.0, 0.0, 1.0])
def load_mesh_dat(path: Path) -> dict:
lines = path.read_text(encoding="utf-8", errors="replace").splitlines()
i = 0
i = read_tag(lines, i, "NbrVertex")
n_vertex = int(lines[i].strip())
i += 1
i = read_tag(lines, i, "Vertex")
vertex = np.zeros((n_vertex, 3))
for r in range(n_vertex):
vertex[r] = np.array(list(map(float, lines[i].split())))
i += 1
i = read_tag(lines, i, "NbrTet")
n_tet = int(lines[i].strip())
i += 1
i = read_tag(lines, i, "Tet")
tet = np.zeros((n_tet, 4), dtype=np.int64)
for r in range(n_tet):
tet[r] = np.array(list(map(int, lines[i].split())), dtype=np.int64)
i += 1
i = read_tag(lines, i, "DomainOfTet")
i += n_tet
i = read_tag(lines, i, "NbrEdge")
n_edge = int(lines[i].strip())
i += 1
i = read_tag(lines, i, "Edge")
i += n_edge
i = read_tag(lines, i, "EdgeOfTet")
i += n_tet
i = read_tag(lines, i, "NbrTri")
n_tri = int(lines[i].strip())
i += 1
i = read_tag(lines, i, "Tri")
tri = np.zeros((n_tri, 3), dtype=np.int64)
for r in range(n_tri):
tri[r] = np.array(list(map(int, lines[i].split())), dtype=np.int64)
i += 1
i = read_tag(lines, i, "DomainOfTri")
domain_of_tri = np.zeros(n_tri, dtype=np.int64)
for r in range(n_tri):
domain_of_tri[r] = int(lines[i].strip())
i += 1
i = read_tag(lines, i, "ConnOfTri")
conn_of_tri = np.zeros((n_tri, 2), dtype=np.int64)
for r in range(n_tri):
conn_of_tri[r] = np.array(list(map(int, lines[i].split())), dtype=np.int64)
i += 1
norm_of_face: dict[int, np.ndarray] = {}
norm_line = i
norm_truncated = False
if i < len(lines) and lines[i].strip() == "NormOfFace":
i += 1
if i >= len(lines):
norm_truncated = True
else:
nbr_norm = int(lines[i].strip())
i += 1
for _ in range(nbr_norm):
if i >= len(lines):
norm_truncated = True
break
parts = list(map(float, lines[i].split()))
if len(parts) < 4:
raise ValueError(f"Bad NormOfFace at line {i + 1}: {lines[i]!r}")
norm_of_face[int(parts[0])] = np.array(parts[1:4])
i += 1
return {
"lines": lines,
"norm_start": norm_line,
"content_end": i,
"norm_truncated": norm_truncated,
"vertex": vertex,
"tet": tet,
"tri": tri,
"domain_of_tri": domain_of_tri,
"conn_of_tri": conn_of_tri,
"norm_of_face": norm_of_face,
"n_vertex": n_vertex,
"n_tet": n_tet,
"n_edge": n_edge,
"n_tri": n_tri,
}
def compute_normals(mesh: dict) -> dict[int, np.ndarray]:
"""Compute NormOfFace for domains 1-9."""
V, T = mesh["vertex"], mesh["tet"]
dom_tri, conn = mesh["domain_of_tri"], mesh["conn_of_tri"]
result: dict[int, np.ndarray] = {}
for domain in ALL_FACE_DOMAINS:
tri_idx = np.where(dom_tri == domain)[0]
if tri_idx.size == 0:
raise ValueError(f"No triangles on domain {domain}")
if domain == 8:
result[domain] = matlab_inc_normal()
continue
if domain in SBC_OUT:
normals = []
for tri in tri_idx:
num_tet = conn[tri, 0] - 1
verts = V[T[num_tet] - 1]
normals.append(matlab_out_normal(domain, verts))
arr = np.vstack(normals)
n = arr.mean(axis=0)
n /= np.linalg.norm(n)
result[domain] = n
continue
# PBC (1,2,4,5) and PMC (6,7): geometric outward normal
geo = []
for tri in tri_idx:
num_tet = conn[tri, 0] - 1
num_face = conn[tri, 1]
pts = face_vertices(V, T, num_tet, num_face)
geo.append(geometric_face_normal(pts))
arr = np.vstack(geo)
n = arr.mean(axis=0)
n /= np.linalg.norm(n)
result[domain] = n
return result
def validate_mesh(mesh: dict, norms: dict[int, np.ndarray]) -> list[str]:
issues: list[str] = []
dom_set = set(int(d) for d in np.unique(mesh["domain_of_tri"]))
if mesh.get("norm_truncated"):
issues.append("CRITICAL: NormOfFace truncated")
if dom_set != set(ALL_FACE_DOMAINS):
issues.append(f"WARNING: DomainOfTri ids {sorted(dom_set)} (expected 1..9)")
extra = set(mesh["norm_of_face"]) - set(ALL_FACE_DOMAINS)
if extra:
issues.append(f"WARNING: obsolete NormOfFace domain ids {sorted(extra)} (e.g. 14 from SBC template)")
missing = set(ALL_FACE_DOMAINS) - set(mesh["norm_of_face"])
if missing:
issues.append(f"WARNING: missing NormOfFace entries for domains {sorted(missing)}")
if 9 in mesh["norm_of_face"] and np.allclose(mesh["norm_of_face"][9], [1, 0, 0], atol=1e-12):
issues.append("CRITICAL: domain 9 was [1,0,0] (SBC template), not assembly_out normal")
if mesh["tet"].min() < 1 or mesh["tet"].max() > mesh["n_vertex"]:
issues.append("CRITICAL: Tet indices out of range")
return issues
def build_norm_section(norms: dict[int, np.ndarray]) -> list[str]:
lines = ["NormOfFace", str(len(ALL_FACE_DOMAINS))]
for d in ALL_FACE_DOMAINS:
n = norms[d]
lines.append(f"{d} {n[0]:.16g} {n[1]:.16g} {n[2]:.16g}")
return lines
def fix_mesh_file(src: Path, dst: Path | None = None) -> int:
dst = dst or src
mesh = load_mesh_dat(src)
norms = compute_normals(mesh)
issues = validate_mesh(mesh, norms)
print(f"=== Validate: {src} ===")
print(f" NbrVertex={mesh['n_vertex']}, NbrTet={mesh['n_tet']}, "
f"NbrEdge={mesh['n_edge']}, NbrTri={mesh['n_tri']}")
print(f" DomainOfTri: {sorted(set(mesh['domain_of_tri'].tolist()))}")
for msg in issues:
print(f" {msg}")
print("\n=== NormOfFace (domains 1-9) ===")
for d in ALL_FACE_DOMAINS:
tag = ""
if d in SBC_OUT:
tag = " [SBC out, assembly_out.m]"
elif d in SBC_INC:
tag = " [SBC inc, assembly_inc.m]"
elif d in {1, 2, 4, 5}:
tag = " [PBC, geometric]"
else:
tag = " [PMC, geometric]"
n = norms[d]
print(f" {d}: ({n[0]:.6g}, {n[1]:.6g}, {n[2]:.6g}){tag}")
out_lines = mesh["lines"][: mesh["norm_start"]]
out_lines.extend(build_norm_section(norms))
dst.write_text("\n".join(out_lines) + "\n", encoding="ascii")
print(f"\nWrote fixed mesh: {dst}")
return 0 if not any("CRITICAL" in x for x in issues) else 1
def main() -> int:
parser = argparse.ArgumentParser(description="Validate and fix PBCmesh.dat NormOfFace (domains 1-9)")
parser.add_argument(
"mesh",
nargs="?",
type=Path,
default=Path(__file__).resolve().parents[2]
/ "三维matlab代码"
/ "matlab 3D一阶基+散射边界条件+单周期边界"
/ "PBCmesh.dat",
)
parser.add_argument("-o", "--output", type=Path, default=None)
parser.add_argument("--copy-to", type=Path, action="append", default=[])
args = parser.parse_args()
if not args.mesh.is_file():
print(f"Mesh not found: {args.mesh}", file=sys.stderr)
return 2
rc = fix_mesh_file(args.mesh, args.output or args.mesh)
out = args.output or args.mesh
for cp in args.copy_to:
cp.parent.mkdir(parents=True, exist_ok=True)
cp.write_text(out.read_text(encoding="ascii"), encoding="ascii")
print(f"Copied to {cp}")
return rc
if __name__ == "__main__":
raise SystemExit(main())