Source code for upxo.pxtal.twinned_simple_3d.cleaning_3d

"""
cleaning_3d.py
==============
Post-twin voxel topology cleaning for the twinned simple 3D pipeline.

Both cleaning stages use fully vectorised local implementations that are
O(1) in grain count (single array pass), replacing the per-grain loop
approach in the generic ``gsdataops.grid_ops`` versions which scale
poorly on large domains.
"""

import numpy as np
from typing import Optional, Dict, Tuple


[docs] class StructureCleaner3D: """ Post-twin voxel topology cleaner for the twinned simple 3D pipeline. Stage 1 -- Vertex spike removal (one global atomic pass, vectorised). Stage 2 -- Lobe splitting via a single cc3d pass on the full array. If defects persist and ``upscale_fallback`` is True the structure is doubled in each spatial dimension (~8x voxel count) and both stages are re-applied. """ __slots__ = ( 'upscale_fallback', 'split_jitter_deg', 'min_clean_voxels', 'do_spike_removal', 'do_lobe_split', 'lgi_clean', 'all_quats_clean', 'twin_role_clean', 'twin_parent_of_clean', 'split_events', 'spike_count', 'n_splits', 'remaining_spikes', 'remaining_multi', 'upscale_applied', '_rng', ) def __init__( self, upscale_fallback: bool = False, split_jitter_deg: float = 0.0, rng_seed: Optional[int] = None, min_clean_voxels: int = 0, do_spike_removal: bool = True, do_lobe_split: bool = True, ): self.upscale_fallback = upscale_fallback self.split_jitter_deg = split_jitter_deg # Independently disable either cleaning stage -- forward-looking # for when the stages themselves gain alternative implementations # (see clean()'s per-stage gating below). Both default on, # matching the previous unconditional behaviour. self.do_spike_removal: bool = bool(do_spike_removal) self.do_lobe_split: bool = bool(do_lobe_split) # Grains smaller than this are skipped for lobe-splitting topology # check. Larger values speed up S30 at the expense of leaving # very small grains unchecked for topological defects. A value of # 0 (default) checks all grains regardless of size. self.min_clean_voxels: int = int(min_clean_voxels) self.lgi_clean: Optional[np.ndarray] = None self.all_quats_clean: Optional[Dict] = None self.twin_role_clean: Optional[Dict] = None self.twin_parent_of_clean: Optional[Dict] = None self.split_events: Dict = {} self.spike_count: int = 0 self.n_splits: int = 0 self.remaining_spikes: int = 0 self.remaining_multi: int = 0 self.upscale_applied: bool = False self._rng = np.random.default_rng(rng_seed) # ── fast private helpers ────────────────────────────────────────────────── def _count_face_same_grain(self, lgi: np.ndarray) -> np.ndarray: """ For every voxel, count how many of its 6 face-connected neighbours carry the same grain label. Single vectorised pass — O(1) in grain count. Returns ------- ndarray, dtype int32, same shape as *lgi*. """ c = np.zeros_like(lgi, dtype=np.int32) for ax in range(3): sl1 = [slice(None)] * 3 sl2 = [slice(None)] * 3 sl1[ax] = slice(1, None) sl2[ax] = slice(None, -1) eq = (lgi[tuple(sl1)] == lgi[tuple(sl2)]).astype(np.int32) c[tuple(sl1)] += eq c[tuple(sl2)] += eq return c # A voxel with same_count <= this many face-connected same-grain # neighbours is treated as a spike. Was same_count == 0 (i.e. <= 0) # until an empirical test against real MC-grown structures showed that # criterion can only ever match fully-isolated singleton voxels, which # are always already consumed by small-grain merging before spike # removal runs -- so it never fired. Loosening to <= 1 catches genuine # thin-necked protrusions (same_count == 1, the neck voxel) without # over-flagging: a same_count == 2 comparison showed those extra # voxels have roughly double the local same-grain mass and typically # touch 2-4 distinct grains -- normal rough/faceted boundary or # triple-junction geometry, not defects -- so <= 2 was rejected. _SPIKE_SAME_COUNT_THRESHOLD = 1 def _fast_clean_spikes(self, lgi: np.ndarray): """ Remove spike voxels (<= _SPIKE_SAME_COUNT_THRESHOLD face-connected same-grain neighbours) in one global atomic pass. Detection is fully vectorised; the repair loop only iterates over detected spikes, which are typically very few. A grain's very last remaining voxel is never reassigned away, even if it's spike-classified -- without this, a thin sliver or single-voxel grain composed entirely of spike voxels could be erased completely in one pass, leaving its grain ID referenced by nothing in `lgi` while every other per-grain array (orientations, twin role, twin parent) still carries an entry for it. `grain_counts` is seeded from `lgi` as it stood before this pass and decremented as this same pass's own reassignments happen, so two spike voxels belonging to a 2-voxel grain still correctly protect only the second one, not both. Parameters ---------- lgi : ndarray (modified in place) Returns ------- lgi, spike_count """ face_offsets = [ (1,0,0),(-1,0,0),(0,1,0),(0,-1,0),(0,0,1),(0,0,-1)] sx, sy, sz = lgi.shape same_count = self._count_face_same_grain(lgi) spike_coords = np.argwhere( (lgi > 0) & (same_count <= self._SPIKE_SAME_COUNT_THRESHOLD)) spike_count = 0 grain_counts = np.bincount(lgi.ravel()) for ix, iy, iz in spike_coords.tolist(): gid = int(lgi[ix, iy, iz]) if grain_counts[gid] <= 1: continue # last voxel of this grain -- never reassign it away nb_counts: Dict[int, int] = {} for dx, dy, dz in face_offsets: nx_, ny_, nz_ = ix+dx, iy+dy, iz+dz if 0 <= nx_ < sx and 0 <= ny_ < sy and 0 <= nz_ < sz: ng = int(lgi[nx_, ny_, nz_]) if ng > 0 and ng != gid: nb_counts[ng] = nb_counts.get(ng, 0) + 1 if nb_counts: lgi[ix, iy, iz] = max(nb_counts, key=nb_counts.get) spike_count += 1 grain_counts[gid] -= 1 return lgi, spike_count def _fast_split_lobes( self, lgi: np.ndarray, all_quats: Dict, twin_role: Dict, twin_parent_of: Dict, next_gid: int, skip_gids: set = None, ): """ Split edge/corner-only connected lobes via a **single** cc3d pass on the full array, rather than one ``ndlabel`` call per grain. Returns ------- lgi, all_quats, twin_role, twin_parent_of, split_events, next_gid """ import cc3d from collections import defaultdict split_events: Dict = {} # Single connected-components pass on the full label array cc = cc3d.connected_components(lgi.astype(np.uint32), connectivity=6) # Vectorised: map each cc component to its original grain ID cc_flat = cc.ravel() lgi_flat = lgi.ravel() valid = cc_flat > 0 cc_valid = cc_flat[valid] lgi_valid = lgi_flat[valid] sort_idx = np.argsort(cc_valid, kind='stable') cc_sorted = cc_valid[sort_idx] lgi_sorted= lgi_valid[sort_idx] first_occ = np.concatenate( [[0], np.where(np.diff(cc_sorted) != 0)[0] + 1]) cc_ids = cc_sorted[first_occ] grain_ids = lgi_sorted[first_occ] # Build grain → [cc component ids] mapping grain_to_ccs: Dict = defaultdict(list) for cc_id, gid in zip(cc_ids.tolist(), grain_ids.tolist()): grain_to_ccs[int(gid)].append(int(cc_id)) _skip = skip_gids or set() for gid, ccs in grain_to_ccs.items(): if len(ccs) <= 1: continue if gid in _skip: continue # small grain — skip lobe check # Sort components by voxel count; largest keeps original ID comp_sizes = sorted( [(cc_id, int(np.sum(cc == cc_id))) for cc_id in ccs], key=lambda x: x[1], reverse=True, ) parent_q = all_quats.get(gid) parent_role = twin_role.get(gid, 'non_host') for rank, (comp_id, comp_size) in enumerate(comp_sizes): if rank == 0: continue lgi[cc == comp_id] = next_gid q_new = parent_q.copy() if parent_q is not None \ else np.array([1., 0., 0., 0.]) jitter_applied = 0.0 if self.split_jitter_deg > 0.0 and parent_q is not None: try: from upxo.xtalphy.crystal_orientation import \ apply_orientation_jitter q_new = apply_orientation_jitter( parent_q, self.split_jitter_deg, self._rng) cos_h = float(np.clip( np.abs(np.dot(parent_q, q_new)), 0., 1.)) jitter_applied = float( np.degrees(2. * np.arccos(cos_h))) except Exception: pass all_quats[next_gid] = q_new twin_role[next_gid] = parent_role twin_parent_of[next_gid] = gid split_events[next_gid] = { 'parent_gid': gid, 'component_size_vox': comp_size, 'jitter_applied_deg': jitter_applied, } next_gid += 1 return lgi, all_quats, twin_role, twin_parent_of, split_events, next_gid def _count_defects(self, lgi: np.ndarray) -> Dict: """Count remaining spikes and multi-component grains.""" import cc3d same_count = self._count_face_same_grain(lgi) spikes = int(np.sum((lgi > 0) & (same_count <= self._SPIKE_SAME_COUNT_THRESHOLD))) cc = cc3d.connected_components(lgi.astype(np.uint32), connectivity=6) cc_flat = cc.ravel() lgi_flat = lgi.ravel() valid = cc_flat > 0 cc_valid = cc_flat[valid] lgi_valid= lgi_flat[valid] sort_idx = np.argsort(cc_valid, kind='stable') lgi_sorted = lgi_valid[sort_idx] cc_sorted = cc_valid[sort_idx] first_occ = np.concatenate( [[0], np.where(np.diff(cc_sorted) != 0)[0] + 1]) grain_ids = lgi_sorted[first_occ] from collections import Counter multi = sum(1 for cnt in Counter(grain_ids.tolist()).values() if cnt > 1) return {'spikes': spikes, 'multi': multi} # ── public API ────────────────────────────────────────────────────────────
[docs] def clean( self, lgi: np.ndarray, all_quats: Dict, twin_role: Dict, twin_parent_of: Dict, ): """ Run Stage 1 (spike removal) + Stage 2 (lobe splitting) on the post-twin grain structure using fast vectorised local methods. Grains with fewer than ``self.min_clean_voxels`` voxels are skipped in Stage 2. Set via the ``min_clean_voxels`` constructor argument. """ import time as _t _t0 = _t.perf_counter() nx, ny, nz = lgi.shape n_grains = int((lgi > 0).sum()) print(f'[S30] StructureCleaner3D domain={nx}x{ny}x{nz} ' f'active_vox={n_grains:,} min_clean_vox={self.min_clean_voxels}') lgi_c = lgi.copy() q_c = dict(all_quats) role_c = dict(twin_role) par_c = dict(twin_parent_of) # Stage 1 — spikes if self.do_spike_removal: print(' [S30] Step 1/2 Spike removal...', end='', flush=True) _t1 = _t.perf_counter() lgi_c, self.spike_count = self._fast_clean_spikes(lgi_c) print(f' done {self.spike_count} spikes ({_t.perf_counter()-_t1:.1f}s)') else: print(' [S30] Step 1/2 Spike removal... SKIPPED (disabled)') self.spike_count = 0 # Stage 2 — lobes (skip grains below min_clean_voxels) if self.do_lobe_split: print(' [S30] Step 2/2 Lobe splitting...', end='', flush=True) _t2 = _t.perf_counter() all_gids = sorted(int(g) for g in np.unique(lgi_c) if g > 0) next_gid = max(all_gids) + 1 if self.min_clean_voxels > 0: # Build size map once with bincount, skip small grains in the lobe pass _counts = np.bincount(lgi_c.ravel(), minlength=next_gid) _skip = {g for g in all_gids if _counts[g] < self.min_clean_voxels} n_skip = len(_skip) else: _skip = set() n_skip = 0 lgi_c, q_c, role_c, par_c, self.split_events, next_gid = \ self._fast_split_lobes(lgi_c, q_c, role_c, par_c, next_gid, skip_gids=_skip) self.n_splits = len(self.split_events) print(f' done {self.n_splits} splits skipped {n_skip} small grains ' f'({_t.perf_counter()-_t2:.1f}s)') else: print(' [S30] Step 2/2 Lobe splitting... SKIPPED (disabled)') self.split_events = {} self.n_splits = 0 rem = self._count_defects(lgi_c) self.remaining_spikes = rem['spikes'] self.remaining_multi = rem['multi'] # Upscale fallback -- only triggered by, and only re-applies, the # stage(s) actually enabled above; a disabled stage's leftover # defects are expected (the user asked to skip it), not a failure # to retry with a bigger domain. trigger_spike = self.do_spike_removal and self.remaining_spikes > 0 trigger_multi = self.do_lobe_split and self.remaining_multi > 0 if (trigger_spike or trigger_multi) and self.upscale_fallback: print('StructureCleaner3D: upscale fallback (2x per axis, ~8x voxels).') print(' WARNING: FEM element count will increase ~8x.') lgi_c = self.apply_upscale(lgi_c, factor=2) self.upscale_applied = True if self.do_spike_removal: lgi_c, sp2 = self._fast_clean_spikes(lgi_c) self.spike_count += sp2 if self.do_lobe_split: all_gids2 = sorted(int(g) for g in np.unique(lgi_c) if g > 0) next_gid2 = max(all_gids2) + 1 lgi_c, q_c, role_c, par_c, ev2, _ = self._fast_split_lobes( lgi_c, q_c, role_c, par_c, next_gid2) self.split_events.update(ev2) self.n_splits = len(self.split_events) rem2 = self._count_defects(lgi_c) self.remaining_spikes = rem2['spikes'] self.remaining_multi = rem2['multi'] self.lgi_clean = lgi_c self.all_quats_clean = q_c self.twin_role_clean = role_c self.twin_parent_of_clean = par_c if self.remaining_spikes == 0 and self.remaining_multi == 0: print(f'StructureCleaner3D: PASS ' f'(spikes removed: {self.spike_count}, ' f'lobes split: {self.n_splits})') else: print(f'StructureCleaner3D: WARNING — ' f'{self.remaining_spikes} spikes, ' f'{self.remaining_multi} multi-component grains remain.')
[docs] def check_remaining_defects(self) -> Dict: return { 'remaining_spikes': self.remaining_spikes, 'remaining_multi': self.remaining_multi, 'upscale_applied': self.upscale_applied, }
[docs] def apply_upscale(self, lgi: np.ndarray, factor: int = 2) -> np.ndarray: result = lgi for ax in range(3): result = np.repeat(result, factor, axis=ax) return result
[docs] def split_events_report(self) -> str: if not self.split_events: return 'No split events.' lines = [f'Split events: {len(self.split_events)}'] for new_gid, info in self.split_events.items(): j = (f', jitter {info["jitter_applied_deg"]:.2f} deg' if info['jitter_applied_deg'] > 0 else ' (same orientation)') lines.append( f' {new_gid} <- parent {info["parent_gid"]} ' f'| {info["component_size_vox"]} vox{j}') return '\n'.join(lines)
[docs] def jitter_report(self) -> str: """Report focused specifically on the orientation jitter applied to split-off lobes -- split_events_report() covers every split event (parent, voxel size, jitter all together in one line); this isolates just the jitter part with summary statistics (min/max/ mean applied jitter), since that's the piece someone tracking orientation-noise/meshing concerns wants on its own, not mixed in with lobe size/parent bookkeeping.""" if not self.split_events: return 'No split events -- no orientation jitter to report.' if self.split_jitter_deg <= 0.0: return (f'Split Jitter was 0 deg for this cleaning pass -- all ' f'{len(self.split_events)} split-off lobe(s) retained ' f'their parent grain\'s orientation exactly (no jitter ' f'applied).') jitters = [info['jitter_applied_deg'] for info in self.split_events.values()] nonzero = [j for j in jitters if j > 0.0] n_zero = len(jitters) - len(nonzero) lines = [ f'Orientation jitter report (Split Jitter setting: ' f'{self.split_jitter_deg:.2f} deg)', f' Split-off lobes: {len(jitters)}', f' Jitter applied (> 0 deg): {len(nonzero)}', f' No jitter applied (parent had no orientation, or jitter ' f'failed): {n_zero}', ] if nonzero: lines.append( f' Applied jitter (deg): min={min(nonzero):.2f} ' f'max={max(nonzero):.2f} mean={sum(nonzero) / len(nonzero):.2f}') lines.append('') lines.append('Per-lobe detail:') for new_gid, info in self.split_events.items(): j = info['jitter_applied_deg'] note = f'{j:.2f} deg' if j > 0 else 'none (same orientation as parent)' lines.append(f' {new_gid} <- parent {info["parent_gid"]}: {note}') return '\n'.join(lines)
[docs] @classmethod def clean_recursive( cls, lgi: np.ndarray, all_quats: Dict, twin_role: Dict, twin_parent_of: Dict, n_passes: int = 1, upscale_fallback: bool = False, split_jitter_deg: float = 0.0, rng_seed: Optional[int] = None, min_clean_voxels: int = 0, do_spike_removal: bool = True, do_lobe_split: bool = True, verbose: bool = True, ) -> Tuple['StructureCleaner3D', int]: """ Repeats :meth:`clean` up to ``n_passes`` times, feeding each pass's cleaned output (lgi_clean/all_quats_clean/twin_role_clean/ twin_parent_of_clean) into the next pass's input, stopping early the moment a pass finds nothing left to fix (no spikes removed, no lobes split) -- a single pass can occasionally leave a spike that only appears after that pass's own lobe-splitting step, which a subsequent pass then catches. Constructs a fresh ``StructureCleaner3D`` for every pass (same constructor arguments each time). The returned cleaner's ``spike_count``/``n_splits``/``split_events`` are overwritten with the cumulative totals/merged history across every pass actually run, so callers see the full picture regardless of how many passes it took. Merging ``split_events`` dicts across passes is safe -- each pass's lobe-splitting numbering starts from that pass's own already-higher array max (it inherits the previous pass's new grain IDs baked into the array), so passes never reuse the same new grain ID for a different lobe. Parameters ---------- n_passes : int Maximum number of cleaning passes. 1 reproduces a single ``clean()`` call exactly. Other parameters : forwarded to the constructor on every pass. Returns ------- (cleaner, n_passes_run) : (StructureCleaner3D, int) cleaner : the final pass's cleaner, with cumulative spike_count/n_splits/split_events across every pass run. n_passes_run : how many passes actually ran (<= n_passes; less if convergence was reached early). """ lgi_in, quats_in = lgi, all_quats role_in, parent_in = twin_role, twin_parent_of total_spike_count = 0 merged_split_events = {} cleaner = None n_passes_run = 0 for it in range(1, n_passes + 1): if n_passes > 1 and verbose: print(f"-- Clean pass {it}/{n_passes} --") cleaner = cls( upscale_fallback=upscale_fallback, split_jitter_deg=split_jitter_deg, rng_seed=rng_seed, min_clean_voxels=min_clean_voxels, do_spike_removal=do_spike_removal, do_lobe_split=do_lobe_split, ) cleaner.clean(lgi_in, quats_in, role_in, parent_in) n_passes_run = it total_spike_count += cleaner.spike_count merged_split_events.update(cleaner.split_events) if n_passes > 1 and cleaner.spike_count == 0 and cleaner.n_splits == 0: if verbose: print(f"Converged after {it} pass(es) -- no further " "spikes or lobe splits found.") break lgi_in, quats_in = cleaner.lgi_clean, cleaner.all_quats_clean role_in, parent_in = cleaner.twin_role_clean, cleaner.twin_parent_of_clean else: if n_passes > 1 and verbose: print(f"Reached the requested {n_passes} pass(es) " "without fully converging.") # Overwrite with cumulative totals across every pass actually run # -- see this method's docstring for why merging is safe. cleaner.spike_count = total_spike_count cleaner.n_splits = len(merged_split_events) cleaner.split_events = merged_split_events return cleaner, n_passes_run return '\n'.join(lines)