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

437 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
"""Plot normE on boundary face domains (e.g. PBC dst faces 2 and 5)."""
from __future__ import annotations
import argparse
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
from matplotlib.collections import PolyCollection
plt.rcParams["font.sans-serif"] = ["Microsoft YaHei", "SimHei", "Arial"]
plt.rcParams["axes.unicode_minus"] = False
# Reuse mesh loader
import importlib.util
_ROOT = Path(__file__).resolve().parents[1]
_spec = importlib.util.spec_from_file_location(
"validate_pbcmesh", _ROOT / "tools" / "validate_and_fix_pbcmesh.py"
)
_v = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(_v)
load_mesh_dat = _v.load_mesh_dat
face_vertices = _v.face_vertices
def face_vertex_ids(tet: np.ndarray, num_tet: int, num_face: int) -> np.ndarray:
"""Return 3 mesh vertex ids (0-based) for a boundary face."""
tv = tet[num_tet - 1] - 1
if num_face == 1:
local = [0, 1, 2]
elif num_face == 2:
local = [0, 1, 3]
elif num_face == 3:
local = [0, 2, 3]
elif num_face == 4:
local = [1, 2, 3]
else:
raise ValueError(f"bad face id {num_face}")
return tv[np.array(local)]
def load_normE(path: Path) -> np.ndarray:
vals: list[float] = []
for line in path.read_text(encoding="utf-8", errors="replace").splitlines():
line = line.strip()
if not line:
continue
if line.startswith("//"):
line = line[2:].strip()
if not line:
continue
vals.append(float(line.split()[0]))
return np.array(vals)
def collect_face_tris(
mesh: dict, domain: int, normE: np.ndarray
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Return (points Nx2 in xz, triangles Mx3, face_values M)."""
V = mesh["vertex"]
T = mesh["tet"]
dom = mesh["domain_of_tri"]
conn = mesh["conn_of_tri"]
pts_list: list[np.ndarray] = []
tri_list: list[tuple[int, int, int]] = []
vals_list: list[float] = []
pt_index: dict[tuple[float, float], int] = {}
def get_pt(x: float, z: float) -> int:
key = (round(x, 10), round(z, 10))
if key not in pt_index:
pt_index[key] = len(pts_list)
pts_list.append([x, z])
return pt_index[key]
for tri_idx in np.where(dom == domain)[0]:
num_tet, num_face = int(conn[tri_idx, 0]), int(conn[tri_idx, 1])
vids = face_vertex_ids(T, num_tet, num_face)
xyz = V[vids]
i0 = get_pt(xyz[0, 0], xyz[0, 2])
i1 = get_pt(xyz[1, 0], xyz[1, 2])
i2 = get_pt(xyz[2, 0], xyz[2, 2])
tri_list.append((i0, i1, i2))
vals_list.append(float(np.mean(normE[vids])))
if not pts_list:
return np.zeros((0, 2)), np.zeros((0, 3), dtype=int), np.zeros(0)
return np.array(pts_list), np.array(tri_list, dtype=int), np.array(vals_list)
def plot_domain_ax(ax, pts: np.ndarray, tris: np.ndarray, fvals: np.ndarray, title: str):
if len(pts) == 0:
ax.set_title(title + " (empty)")
return
# Per-triangle constant color
polys = pts[tris]
coll = PolyCollection(
polys,
array=fvals,
cmap="jet",
edgecolors="k",
linewidths=0.15,
alpha=0.95,
)
ax.add_collection(coll)
ax.set_aspect("equal")
ax.set_xlabel("x (m)")
ax.set_ylabel("z (m)")
ax.set_title(title)
ax.autoscale()
mappable = coll
return mappable
def _tri_centroids(pts: np.ndarray, tris: np.ndarray) -> np.ndarray:
return np.array([pts[t].mean(axis=0) for t in tris])
def _centroids_inside_patch(
centroids: np.ndarray, patch_pts: np.ndarray, patch_tris: np.ndarray
) -> np.ndarray:
"""True where centroid lies inside the xz footprint of a patch (bbox test)."""
xmin, xmax = patch_pts[:, 0].min(), patch_pts[:, 0].max()
zmin, zmax = patch_pts[:, 1].min(), patch_pts[:, 1].max()
return (
(centroids[:, 0] >= xmin)
& (centroids[:, 0] <= xmax)
& (centroids[:, 1] >= zmin)
& (centroids[:, 1] <= zmax)
)
def _append_tris_unified(
pts_acc: list[np.ndarray],
tris_acc: list[tuple[int, int, int]],
vals_acc: list[float],
pt_index: dict[tuple[float, float], int],
patch_pts: np.ndarray,
patch_tris: np.ndarray,
patch_vals: np.ndarray,
) -> None:
def get_pt(x: float, z: float) -> int:
key = (round(x, 10), round(z, 10))
if key not in pt_index:
pt_index[key] = len(pts_acc)
pts_acc.append([x, z])
return pt_index[key]
for tri, val in zip(patch_tris, patch_vals):
p0, p1, p2 = patch_pts[tri]
i0 = get_pt(p0[0], p0[1])
i1 = get_pt(p1[0], p1[1])
i2 = get_pt(p2[0], p2[1])
tris_acc.append((i0, i1, i2))
vals_acc.append(float(val))
def collect_domains_combined(
mesh: dict,
domains: list[int],
normE: np.ndarray,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Merge face domains on one x-z plane with shared vertex indexing.
Adjacent patches (e.g. domain 2 + 5) share boundary vertices; duplicate
points at the same (x, z) are welded so triangles tile without z-fighting.
"""
if not domains:
return np.zeros((0, 2)), np.zeros((0, 3), dtype=int), np.zeros(0)
raw = [(d, *collect_face_tris(mesh, d, normE)) for d in domains]
raw = [(d, p, t, f) for d, p, t, f in raw if len(t) > 0]
if not raw:
return np.zeros((0, 2)), np.zeros((0, 3), dtype=int), np.zeros(0)
pts_acc: list[list[float]] = []
tris_acc: list[tuple[int, int, int]] = []
vals_acc: list[float] = []
pt_index: dict[tuple[float, float], int] = {}
for _, pts, tris, fvals in raw:
_append_tris_unified(
pts_acc, tris_acc, vals_acc, pt_index, pts, tris, fvals,
)
if not pts_acc:
return np.zeros((0, 2)), np.zeros((0, 3), dtype=int), np.zeros(0)
return (
np.array(pts_acc),
np.array(tris_acc, dtype=int),
np.array(vals_acc),
)
def plot_combined_faces(
mesh: dict,
domains: list[int],
normE_cpp: np.ndarray,
normE_mat: np.ndarray | None,
output: Path,
title: str | None = None,
cpp_outdir: Path | None = None,
mat_outdir: Path | None = None,
layout: str = "merge",
) -> None:
"""Plot PBC face domains on one axes.
layout:
merge weld shared vertices and tile all domains (default)
side one panel per domain, horizontal row
"""
has_mat = normE_mat is not None
plot_domains = list(domains)
use_side = layout == "side" and len(plot_domains) > 1
ncols = (len(plot_domains) if use_side else 1) * (2 if has_mat else 1)
fig_w = 5.5 * ncols if use_side else (6.8 * ncols if has_mat else 6.0)
fig, axes = plt.subplots(
1, ncols if ncols else 1,
figsize=(fig_w, 5.5),
squeeze=False,
facecolor="white",
)
cpp_is_mat = (
has_mat
and cpp_outdir is not None
and mat_outdir is not None
and cpp_outdir.resolve() == mat_outdir.resolve()
)
mappables = []
labels = ["C++ OpticsFEM", "MATLAB 参考"] if has_mat else ["OpticsFEM |E|"]
if cpp_is_mat:
labels[0] = "MATLAB 参考 (C++ OutFile 待 pbc3d_sbc.json 重跑)"
datasets = [normE_cpp]
if has_mat:
datasets.append(normE_mat)
domain_label = " + ".join(f"domain {d}" for d in plot_domains)
if use_side:
mappables = []
for i, dom in enumerate(plot_domains):
col = i * (2 if has_mat else 1)
for j, (normE, suffix) in enumerate(
[(normE_cpp, "C++"), (normE_mat, "MATLAB")] if has_mat else [(normE_cpp, "")]
):
if normE is None:
continue
pts, tris, fvals = collect_face_tris(mesh, dom, normE)
ax = axes[0, col + j]
label = f"{suffix} domain {dom}" if suffix else f"domain {dom}"
m = plot_domain_ax(ax, pts, tris, fvals, f"{label}\n(n_tri={len(fvals)})")
if m is not None:
mappables.append(m)
else:
for col, (normE, label) in enumerate(zip(datasets, labels)):
pts, tris, fvals = collect_domains_combined(mesh, plot_domains, normE)
n_tri = len(fvals)
m = plot_domain_ax(
axes[0, col],
pts,
tris,
fvals,
f"{label}\n(domain {' + '.join(map(str, plot_domains))}, n_tri={n_tri})",
)
if m is not None:
mappables.append(m)
if mappables:
vmin = min(m.get_array().min() for m in mappables if m.get_array() is not None)
vmax = max(m.get_array().max() for m in mappables if m.get_array() is not None)
for m in mappables:
m.set_clim(vmin, vmax)
if has_mat:
fig.colorbar(
mappables[0],
ax=axes.ravel().tolist(),
shrink=0.82,
pad=0.02,
label="|E| (V/m)",
)
else:
fig.subplots_adjust(right=0.88)
cbar = fig.colorbar(
mappables[0],
ax=axes[0, 0],
fraction=0.046,
pad=0.04,
label="|E| (V/m)",
)
cbar.ax.tick_params(labelsize=9)
fig.subplots_adjust(top=0.92, bottom=0.12, left=0.10, right=0.88 if not has_mat else 0.95)
output.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(output, dpi=160, bbox_inches="tight", facecolor="white")
plt.close(fig)
def main() -> None:
parser = argparse.ArgumentParser(description="Plot normE on mesh face domains")
parser.add_argument(
"--mesh",
type=Path,
default=_ROOT / "build" / "Release" / "PBCmesh.dat",
)
parser.add_argument(
"--outdir",
type=Path,
default=_ROOT / "build" / "Release" / "OutFile",
help="OpticsFEM OutFile with normE",
)
parser.add_argument(
"--mat-outdir",
type=Path,
default=_ROOT.parent
/ "三维matlab代码"
/ "matlab 3D一阶基+散射边界条件+单周期边界"
/ "OutFile",
help="MATLAB OutFile for comparison (set empty to skip)",
)
parser.add_argument(
"--domains",
type=int,
nargs="+",
default=[2, 5],
help="DomainOfTri ids to plot",
)
parser.add_argument(
"--output",
type=Path,
default=_ROOT / "build" / "Release" / "OutFile" / "faces_2_5.png",
)
parser.add_argument(
"--combined",
action="store_true",
help="Plot all --domains on the same axes (connected geometry)",
)
parser.add_argument(
"--no-compare",
action="store_true",
help="Skip MATLAB comparison panel (simulation result only)",
)
parser.add_argument(
"--layout",
choices=["merge", "side"],
default="merge",
help="merge: weld vertices and tile domains; side: horizontal panels",
)
args = parser.parse_args()
mesh = load_mesh_dat(args.mesh)
normE_cpp = load_normE(args.outdir / "normE")
if normE_cpp.size != mesh["n_vertex"]:
raise SystemExit(
f"normE size {normE_cpp.size} != NbrVertex {mesh['n_vertex']}"
)
has_mat = (
not args.no_compare
and args.mat_outdir is not None
and (args.mat_outdir / "normE").exists()
)
normE_mat = load_normE(args.mat_outdir / "normE") if has_mat else None
if args.combined:
plot_combined_faces(
mesh,
args.domains,
normE_cpp,
normE_mat,
args.output,
cpp_outdir=args.outdir,
mat_outdir=args.mat_outdir if has_mat else None,
layout=args.layout,
)
print(f"Saved combined {args.output}")
return
ncols = 2 if has_mat else 1
nrows = len(args.domains)
fig, axes = plt.subplots(
nrows, ncols, figsize=(6.5 * ncols, 5 * nrows), squeeze=False, facecolor="white"
)
mappables = []
for row, dom in enumerate(args.domains):
pts, tris, fvals = collect_face_tris(mesh, dom, normE_cpp)
m = plot_domain_ax(
axes[row, 0],
pts,
tris,
fvals,
f"C++ domain {dom} (n_tri={len(fvals)})",
)
if m is not None:
mappables.append(m)
if has_mat:
pts_m, tris_m, fvals_m = collect_face_tris(mesh, dom, normE_mat)
m2 = plot_domain_ax(
axes[row, 1],
pts_m,
tris_m,
fvals_m,
f"MATLAB domain {dom}",
)
if m2 is not None:
mappables.append(m2)
if mappables:
vmin = min(m.get_array().min() for m in mappables if m.get_array() is not None)
vmax = max(m.get_array().max() for m in mappables if m.get_array() is not None)
for m in mappables:
m.set_clim(vmin, vmax)
fig.colorbar(mappables[0], ax=axes.ravel().tolist(), shrink=0.6, label="|E| (V/m)")
fig.suptitle(
"normE on PBC faces (xz projection, y≈0)\n"
+ ", ".join(f"domain {d}" for d in args.domains),
fontsize=13,
fontweight="bold",
)
fig.tight_layout()
args.output.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(args.output, dpi=160, bbox_inches="tight")
print(f"Saved {args.output}")
plt.close(fig)
if __name__ == "__main__":
main()