BIMStruct3D-segmentation / pointcept /models /point_prompt_training /point_prompt_training_v1m3_neo.py
Download pointcept/models/point_prompt_training/point_prompt_training_v1m3_neo.py from dfki-av/BIMStruct3D-segmentation: direct link, hf CLI and curl.
- Browser
- Download file 5.53 kB
-
https://huggingface.co/dfki-av/BIMStruct3D-segmentation/resolve/7ab05dd198d38ea08e60bb1e4f3f04aa369975c2/pointcept/models/point_prompt_training/point_prompt_training_v1m3_neo.py
- Command line
-
hf download hf://dfki-av/BIMStruct3D-segmentation@7ab05dd198d38ea08e60bb1e4f3f04aa369975c2/pointcept/models/point_prompt_training/point_prompt_training_v1m3_neo.py
-
curl -L -o point_prompt_training_v1m3_neo.py https://huggingface.co/dfki-av/BIMStruct3D-segmentation/resolve/7ab05dd198d38ea08e60bb1e4f3f04aa369975c2/pointcept/models/point_prompt_training/point_prompt_training_v1m3_neo.py
5.53 kB
| """ | |
| Point Prompt Training specific for Sonata | |
| In Sonata, as we identify the criminal that restricts domain adaptive capacity is BN | |
| and successfully replaces all BN in PTv3 to LN. I think we don't need PDNorm | |
| anymore. Hence, I remove the domain prompting and turns it into a pure multi-dataset | |
| joint training framework. | |
| Author: Xiaoyang Wu (xiaoyang.wu.cs@gmail.com) | |
| Please cite our work if the code is helpful to you. | |
| """ | |
| from functools import partial | |
| from collections import OrderedDict | |
| import torch | |
| import torch.nn as nn | |
| from pointcept.models.utils.structure import Point | |
| from pointcept.models.builder import MODELS | |
| from pointcept.models.losses import build_criteria | |
| class PointPromptTraining(nn.Module): | |
| """ | |
| PointPromptTraining provides Data-driven Context and enables multi-dataset training with | |
| Language-driven Categorical Alignment. PDNorm is supported by SpUNet-v1m3 to adapt the | |
| backbone to a specific dataset with a given dataset condition and context. | |
| """ | |
| def __init__( | |
| self, | |
| backbone=None, | |
| criteria=None, | |
| backbone_out_channels=96, | |
| conditions=("Structured3D", "ScanNet", "S3DIS"), | |
| template="[x]", | |
| clip_model="ViT-B/16", | |
| class_names=None, | |
| freeze_backbone=False, | |
| backbone_mode=False, | |
| ): | |
| super().__init__() | |
| if class_names is None: | |
| # fmt: off | |
| class_names = [ | |
| ("wall", "floor", "cabinet", "bed", "chair", | |
| "sofa", "table", "door", "window", "picture", | |
| "desk", "shelves", "curtain", "dresser", "pillow", | |
| "mirror", "ceiling", "refrigerator", "television", "nightstand", | |
| "sink", "lamp", "otherstructure", "otherfurniture", "otherprop"), | |
| ("wall", "floor", "cabinet", "bed", "chair", | |
| "sofa", "table", "door", "window", "bookshelf", | |
| "picture", "counter", "desk", "curtain", "refridgerator", | |
| "shower curtain", "toilet", "sink", "bathtub", "otherfurniture"), | |
| ("ceiling", "floor", "wall", "beam", "column", | |
| "window", "door", "table", "chair", "sofa", | |
| "bookcase", "board", "clutter"), | |
| ] | |
| # fmt: on | |
| assert len(conditions) == len(class_names) | |
| # assert backbone.type in ["SpUNet-v1m3", "PT-v2m3", "PT-v3m1"] | |
| self.backbone = MODELS.build(backbone) | |
| self.criteria = build_criteria(criteria) | |
| self.conditions = conditions | |
| self.freeze_backbone = freeze_backbone | |
| if self.freeze_backbone: | |
| for p in self.backbone.parameters(): | |
| p.requires_grad = False | |
| self.backbone_mode = backbone_mode | |
| if not self.backbone_mode: | |
| import clip | |
| clip_model, _ = clip.load( | |
| clip_model, device="cpu", download_root="./.cache/clip" | |
| ) | |
| clip_model.requires_grad_(False) | |
| class_embeddings = [] | |
| num_classes = [] | |
| for i in range(len(conditions)): | |
| class_prompt = [ | |
| template.replace("[x]", name) for name in class_names[i] | |
| ] | |
| class_token = clip.tokenize(class_prompt) | |
| class_embedding = clip_model.encode_text(class_token) | |
| class_embedding = class_embedding / class_embedding.norm( | |
| dim=-1, keepdim=True | |
| ) | |
| class_embeddings.append(class_embedding) | |
| num_classes.append(len(class_prompt)) | |
| class_embeddings = torch.cat(class_embeddings, dim=0) | |
| self.register_buffer("class_embeddings", class_embeddings, persistent=False) | |
| self.num_classes = num_classes | |
| self.proj_head = nn.Linear( | |
| backbone_out_channels, clip_model.text_projection.shape[1] | |
| ) | |
| self.logit_scale = clip_model.logit_scale | |
| def forward(self, data_dict): | |
| condition = data_dict["condition"][0] | |
| if self.freeze_backbone: | |
| with torch.no_grad(): | |
| point = self.backbone(data_dict) | |
| else: | |
| point = self.backbone(data_dict) | |
| while "pooling_parent" in point.keys(): | |
| assert "pooling_inverse" in point.keys() | |
| parent = point.pop("pooling_parent") | |
| inverse = point.pop("pooling_inverse") | |
| parent.feat = torch.cat([parent.feat, point.feat[inverse]], dim=-1) | |
| point = parent | |
| if self.backbone_mode: | |
| # PPT serve as a multi-dataset backbone when enable backbone mode | |
| return point.feat | |
| feat = self.proj_head(point.feat) | |
| eps = 1e-6 if feat.dtype == torch.float16 else 1e-12 | |
| feat = nn.functional.normalize(feat, dim=-1, p=2, eps=eps) | |
| sim = ( | |
| feat | |
| self.conditions.index(condition) | |
| ].t() | |
| ) | |
| logit_scale = self.logit_scale.exp() | |
| seg_logits = logit_scale * sim | |
| # train | |
| if self.training: | |
| loss = self.criteria(seg_logits, data_dict["segment"]) | |
| return dict(loss=loss) | |
| # eval | |
| elif "segment" in data_dict.keys(): | |
| loss = self.criteria(seg_logits, data_dict["segment"]) | |
| return dict(loss=loss, seg_logits=seg_logits) | |
| # test | |
| else: | |
| return dict(seg_logits=seg_logits) | |