Source code for pimm.datasets.transform.multiview

"""Contrastive and multiview SSL transform builders."""

from .common import *
from .base import Compose


[docs] @TRANSFORMS.register_module() class ContrastiveViewsGenerator(object): """Generate two independently-augmented views for contrastive SSL. Copies the keys in ``view_keys`` into two sub-dicts, runs the same transform pipeline (built from ``view_trans_cfg``) independently on each, then writes the results back under ``view1_<key>`` and ``view2_<key>`` prefixes (for every key produced by the pipeline). Registered as ``ContrastiveViewsGenerator`` — use this string as the ``type`` in a ``transform=[...]`` config list. Args: view_keys (tuple): Keys copied from the source sample into each view. Defaults to ``("coord", "color", "normal", "origin_coord")``. view_trans_cfg (list, optional): Transform config dicts composed into the per-view pipeline applied to both views. Defaults to ``None``. Example: .. code-block:: python >>> import numpy as np >>> np.random.seed(0) >>> data = {"coord": np.random.rand(10, 3).astype("f4"), ... "color": np.random.rand(10, 3).astype("f4")} >>> out = ContrastiveViewsGenerator( ... view_keys=("coord", "color"), ... view_trans_cfg=[dict(type="RandomFlip", p=1.0)])(data) >>> sorted(k for k in out if k.startswith("view")) ['view1_color', 'view1_coord', 'view2_color', 'view2_coord'] # two independently-augmented copies of each view_key """ def __init__( self, view_keys=("coord", "color", "normal", "origin_coord"), view_trans_cfg=None, ): self.view_keys = view_keys self.view_trans = Compose(view_trans_cfg) def __call__(self, data_dict): view1_dict = dict() view2_dict = dict() for key in self.view_keys: view1_dict[key] = data_dict[key].copy() view2_dict[key] = data_dict[key].copy() view1_dict = self.view_trans(view1_dict) view2_dict = self.view_trans(view2_dict) for key, value in view1_dict.items(): data_dict["view1_" + key] = value for key, value in view2_dict.items(): data_dict["view2_" + key] = value return data_dict
[docs] @TRANSFORMS.register_module() class MultiViewGenerator(object): """Generate multiple global and local crops for DINO/iBOT-style SSL. Crops a major global view around a randomly chosen center (restricted in height by ``center_height_scale``), then additional global views and several local views, each a nearest-``k`` neighborhood whose size is a random fraction of the cloud (per ``global_view_scale`` / ``local_view_scale``). Local centers are pushed to cover the major view and, when ``data_dict["anchors"]`` is available, a fraction (``anchor_bias_ratio``) of local crops are centered on anchors. Reads ``data_dict["coord"]`` and the keys in ``view_keys`` (and optional ``anchors``); writes concatenated ``global_<key>`` / ``local_<key>`` arrays with ``global_offset`` / ``local_offset`` boundaries. Registered as ``MultiViewGenerator`` — use this string as the ``type`` in a ``transform=[...]`` config list. Args: global_view_num (int): Number of global views (including the major one). Defaults to ``2``. global_view_scale (tuple): ``(min, max)`` fraction of points per global crop. Defaults to ``(0.4, 1.0)``. local_view_num (int): Number of local views. Defaults to ``4``. local_view_scale (tuple): ``(min, max)`` fraction of points per local crop. Defaults to ``(0.1, 0.4)``. global_shared_transform (list, optional): Transform config applied once to the whole sample before any cropping. Defaults to ``None``. global_transform (list, optional): Transform config applied to each global view. Defaults to ``None``. local_transform (list, optional): Transform config applied to each local view. Defaults to ``None``. max_size (int): Upper bound on the points considered per view. Defaults to ``65536``. center_height_scale (tuple): ``(min, max)`` fractional z-band restricting where the major-view center is sampled. Defaults to ``(0, 1)``. shared_global_view (bool): If ``True``, the extra global views are copies of the major view rather than freshly cropped. Defaults to ``False``. center_sampling (str): Center-selection strategy, ``"random"`` or ``"cnms"`` (coverage non-max suppression). Defaults to ``"random"``. center_sampling_kwargs (dict, optional): Extra kwargs for ``cnms`` when ``center_sampling="cnms"``. Defaults to ``None``. view_keys (tuple): Keys cropped into each view; must include ``"coord"``. Defaults to ``("coord", "origin_coord", "color", "normal")``. anchor_bias_ratio (float): Fraction of local views centered on anchors (when anchors are present). Defaults to ``0.6``. anchor_radius_scale (float): Radius scale for anchor-centered local crops (size scales roughly cubically with it). Defaults to ``1.5``. anchor_keys (tuple): Anchor sub-keys pooled as candidate centers; ``led`` is always excluded. Defaults to ``("endpoints", "branches_track", "branches_shower", "bragg")``. Example: .. code-block:: python >>> import numpy as np >>> np.random.seed(0) >>> data = {"coord": np.random.rand(200, 3).astype("f4"), ... "color": np.random.rand(200, 3).astype("f4"), ... "index_valid_keys": ["coord", "color"]} >>> out = MultiViewGenerator(global_view_num=2, local_view_num=4, ... view_keys=("coord", "color"))(data) >>> out["global_coord"].shape, out["global_offset"] # 2 global crops concatenated ((356, 3), array([170, 356])) >>> out["local_coord"].shape, out["local_offset"] # 4 local crops concatenated ((223, 3), array([ 66, 134, 155, 223])) # exact sizes vary with the RNG; offsets mark per-view boundaries """ def __init__( self, global_view_num=2, global_view_scale=(0.4, 1.0), local_view_num=4, local_view_scale=(0.1, 0.4), global_shared_transform=None, global_transform=None, local_transform=None, max_size=65536, center_height_scale=(0, 1), shared_global_view=False, center_sampling="random", # or cnms center_sampling_kwargs=None, view_keys=("coord", "origin_coord", "color", "normal"), # Anchor-biased sampling anchor_bias_ratio=0.6, anchor_radius_scale=1.5, anchor_keys=("endpoints", "branches_track", "branches_shower", "bragg"), ): self.global_view_num = global_view_num self.global_view_scale = global_view_scale self.local_view_num = local_view_num self.local_view_scale = local_view_scale self.global_shared_transform = Compose(global_shared_transform) self.global_transform = Compose(global_transform) self.local_transform = Compose(local_transform) self.max_size = max_size self.center_height_scale = center_height_scale self.shared_global_view = shared_global_view self.view_keys = view_keys assert "coord" in view_keys self.center_sampling = center_sampling self.center_sampling_kwargs = center_sampling_kwargs # Anchors self.anchor_bias_ratio = anchor_bias_ratio self.anchor_radius_scale = anchor_radius_scale self.anchor_keys = anchor_keys
[docs] def get_view(self, point, center, scale, size_override: Optional[int] = None): coord = point["coord"] max_size = min(self.max_size, coord.shape[0]) if max_size <= 0: raise ValueError("Cannot generate a view from an empty point cloud") if size_override is None: size = int(np.random.uniform(*scale) * max_size) else: size = int(size_override) size = max(1, min(max_size, size)) index = np.argsort(np.sum(np.square(coord - center), axis=-1))[:size] view = dict(index=index) for key in point.keys(): if key in self.view_keys: view[key] = point[key][index] if "index_valid_keys" in point.keys(): # inherit index_valid_keys from point view["index_valid_keys"] = point["index_valid_keys"] return view
[docs] def get_center(self, coord, mask=None): if mask is None: possible_centers = coord else: possible_centers = coord[np.where(mask)[0]] if self.center_sampling == "cnms": from cnms import cnms possible_centers, _, _ = cnms(possible_centers, **self.center_sampling_kwargs) return possible_centers[np.random.choice(possible_centers.shape[0])]
def _build_global_views(self, point, major_view): major_coord = major_view["coord"] if not self.shared_global_view: global_views = [ self.get_view( point=point, center=major_coord[np.random.randint(major_coord.shape[0])], scale=self.global_view_scale, ) for _ in range(self.global_view_num - 1) ] else: global_views = [ {key: value.copy() for key, value in major_view.items()} for _ in range(self.global_view_num - 1) ] return [major_view] + global_views def _pack_views(self, data_dict, global_views, local_views): view_dict = {} for global_view in global_views: global_view.pop("index") global_view = self.global_transform(global_view) for key in self.view_keys: if f"global_{key}" in view_dict.keys(): view_dict[f"global_{key}"].append(global_view[key]) else: view_dict[f"global_{key}"] = [global_view[key]] view_dict["global_offset"] = np.cumsum( [data.shape[0] for data in view_dict["global_coord"]] ) for local_view in local_views: local_view.pop("index") local_view = self.local_transform(local_view) for key in self.view_keys: if f"local_{key}" in view_dict.keys(): view_dict[f"local_{key}"].append(local_view[key]) else: view_dict[f"local_{key}"] = [local_view[key]] view_dict["local_offset"] = np.cumsum( [data.shape[0] for data in view_dict["local_coord"]] ) for key in view_dict.keys(): if "offset" not in key: view_dict[key] = np.concatenate(view_dict[key], axis=0) data_dict.update(view_dict) return data_dict def __call__(self, data_dict): coord = data_dict["coord"] point = self.global_shared_transform(copy.deepcopy(data_dict)) z_min = coord[:, 2].min() z_max = coord[:, 2].max() z_min_ = z_min + (z_max - z_min) * self.center_height_scale[0] z_max_ = z_min + (z_max - z_min) * self.center_height_scale[1] center_mask = np.logical_and(coord[:, 2] >= z_min_, coord[:, 2] <= z_max_) # get major global view major_center = coord[np.random.choice(np.where(center_mask)[0])] major_view = self.get_view(point, major_center, self.global_view_scale) major_coord = major_view["coord"] # get global views: restrict the center of left global view within the major global view global_views = self._build_global_views(point, major_view) # get local views: restrict the center of local view within the major global view cover_mask = np.zeros_like(major_view["index"], dtype=bool) local_views = [] # Prepare anchor pool if available (exclude LEDs) anchors_pool = [] if isinstance(data_dict.get("anchors"), dict): for k in self.anchor_keys: if k == "led": continue v = data_dict["anchors"].get(k) if v is not None and len(v) > 0: anchors_pool.append(v) anchors_pool = np.concatenate(anchors_pool, axis=0) if len(anchors_pool) > 0 else np.zeros((0,3), dtype=np.float32) # Map anchors to nearest point inside major view to keep locality consistent kd_major = cKDTree(major_coord) if major_coord.shape[0] > 0 else None # Estimate size override for anchor crops: approximate radius scaling via cubic relation # size' ~= size * (radius_scale^3) size_base = int(np.mean([np.random.uniform(*self.local_view_scale) * min(self.max_size, coord.shape[0]) for _ in range(4)])) size_override = int(max(8, min(self.max_size, size_base * (self.anchor_radius_scale ** 3)))) # Determine counts num_anchor_locals = int(np.ceil(self.local_view_num * float(self.anchor_bias_ratio))) if anchors_pool.shape[0] > 0 else 0 num_random_locals = self.local_view_num - num_anchor_locals # Anchor-centered locals for i in range(num_anchor_locals): if sum(~cover_mask) == 0: cover_mask[:] = False if anchors_pool.shape[0] == 0: break aidx = np.random.randint(0, anchors_pool.shape[0]) acoord = anchors_pool[aidx] # Project to nearest major point to keep within major global view if kd_major is not None and kd_major.n > 0: _, nn = kd_major.query(acoord, k=1) center = major_coord[nn] else: center = acoord local_view = self.get_view( point=data_dict, center=center, scale=self.local_view_scale, size_override=size_override, ) local_views.append(local_view) cover_mask[np.isin(major_view["index"], local_view["index"])] = True # Uniform random locals for i in range(num_random_locals): if sum(~cover_mask) == 0: cover_mask[:] = False local_view = self.get_view( point=data_dict, center=major_coord[np.random.choice(np.where(~cover_mask)[0])], scale=self.local_view_scale, ) local_views.append(local_view) cover_mask[np.isin(major_view["index"], local_view["index"])] = True return self._pack_views(data_dict, global_views, local_views)
[docs] @TRANSFORMS.register_module() class MixedScaleGeometryMultiViewGenerator(MultiViewGenerator): """Multi-view generator with coarse locals plus extra fine local crops. Subclass of :class:`MultiViewGenerator` that replaces ``fine_local_view_num`` of the local crops with much smaller (``fine_local_view_scale``) crops whose centers come either from uniform sampling (``fine_center_mode="random"``) or from a local-PCA directional-complexity score (``"geometry"``, keeping the top ``fine_center_top_frac`` of points by complexity). Keeps the SSL objective unchanged while steering which local regions feed the local-global loss; reads/writes the same keys as the base class. All base-class arguments are accepted via ``**kwargs``. Registered as ``MixedScaleGeometryMultiViewGenerator`` — use this string as the ``type`` in a ``transform=[...]`` config list. Args: fine_local_view_num (int): Number of fine local crops (must satisfy ``0 <= fine_local_view_num <= local_view_num``). Defaults to ``3``. fine_local_view_scale (tuple): ``(min, max)`` fraction of points per fine local crop. Defaults to ``(0.01, 0.04)``. fine_center_mode (str): ``"geometry"`` to bias fine centers toward high-complexity points, or ``"random"`` for uniform. Defaults to ``"geometry"``. fine_center_top_frac (float): Top fraction of points (by complexity) eligible as fine centers in ``"geometry"`` mode. Defaults to ``0.05``. fine_center_k (int): Neighbor count for the local-PCA complexity score. Defaults to ``24``. **kwargs: Forwarded to :class:`MultiViewGenerator`. Example: .. code-block:: python >>> import numpy as np >>> np.random.seed(0) >>> data = {"coord": np.random.rand(300, 3).astype("f4"), ... "energy": np.random.rand(300, 1).astype("f4"), ... "index_valid_keys": ["coord", "energy"]} >>> out = MixedScaleGeometryMultiViewGenerator( ... fine_local_view_num=2, local_view_num=4, ... view_keys=("coord", "energy"))(data) >>> out["local_coord"].shape, out["local_offset"] ((193, 3), array([ 8, 14, 105, 193])) # first 2 locals are tiny geometry-biased fine crops (8, 6 pts), then coarse locals """ def __init__( self, fine_local_view_num=3, fine_local_view_scale=(0.01, 0.04), fine_center_mode="geometry", fine_center_top_frac=0.05, fine_center_k=24, **kwargs, ): super().__init__(**kwargs) assert 0 <= fine_local_view_num <= self.local_view_num assert fine_center_mode in ("geometry", "random") self.fine_local_view_num = int(fine_local_view_num) self.fine_local_view_scale = fine_local_view_scale self.fine_center_mode = fine_center_mode self.fine_center_top_frac = float(fine_center_top_frac) self.fine_center_k = int(fine_center_k) @staticmethod def _directional_complexity(coord, k): coord = np.asarray(coord, dtype=np.float32) n = coord.shape[0] if n < 4: return np.zeros(n, dtype=np.float32) k_eff = min(int(k) + 1, n) tree = cKDTree(coord) try: _, idx = tree.query(coord, k=k_eff, workers=-1) except TypeError: _, idx = tree.query(coord, k=k_eff) if idx.ndim == 1: idx = idx[:, None] idx = idx[:, 1:] if idx.shape[1] < 3: return np.zeros(n, dtype=np.float32) neigh = coord[idx] centered = neigh - neigh.mean(axis=1, keepdims=True) cov = np.einsum("nki,nkj->nij", centered, centered) / centered.shape[1] eig = np.maximum(np.linalg.eigvalsh(cov), 0.0) return (eig[:, 1] / (eig[:, 2] + 1.0e-8)).astype(np.float32) def _geometry_pool(self, coord, major_index): if self.fine_center_mode == "random": return major_index score = self._directional_complexity(coord, self.fine_center_k) n_top = max(1, int(np.ceil(score.shape[0] * self.fine_center_top_frac))) top_index = np.argpartition(score, -n_top)[-n_top:] in_major = np.zeros(coord.shape[0], dtype=bool) in_major[major_index] = True pool = top_index[in_major[top_index]] return pool if pool.shape[0] > 0 else major_index def __call__(self, data_dict): coord = data_dict["coord"] point = self.global_shared_transform(copy.deepcopy(data_dict)) z_min = coord[:, 2].min() z_max = coord[:, 2].max() z_min_ = z_min + (z_max - z_min) * self.center_height_scale[0] z_max_ = z_min + (z_max - z_min) * self.center_height_scale[1] center_mask = np.logical_and(coord[:, 2] >= z_min_, coord[:, 2] <= z_max_) major_center = coord[np.random.choice(np.where(center_mask)[0])] major_view = self.get_view(point, major_center, self.global_view_scale) major_coord = major_view["coord"] global_views = self._build_global_views(point, major_view) cover_mask = np.zeros_like(major_view["index"], dtype=bool) local_views = [] fine_pool = self._geometry_pool(coord, major_view["index"]) for _ in range(self.fine_local_view_num): center = coord[fine_pool[np.random.randint(fine_pool.shape[0])]] local_views.append( self.get_view( point=data_dict, center=center, scale=self.fine_local_view_scale, ) ) num_random_locals = self.local_view_num - self.fine_local_view_num for _ in range(num_random_locals): if sum(~cover_mask) == 0: cover_mask[:] = False local_view = self.get_view( point=data_dict, center=major_coord[np.random.choice(np.where(~cover_mask)[0])], scale=self.local_view_scale, ) local_views.append(local_view) cover_mask[np.isin(major_view["index"], local_view["index"])] = True return self._pack_views(data_dict, global_views, local_views)