Classes Fixed abnd debugging

This commit is contained in:
Si11ium
2020-07-03 14:40:28 +02:00
parent e9d0591b11
commit 5353220890
10 changed files with 66 additions and 59 deletions
-1
View File
@@ -2,7 +2,6 @@ from abc import ABC
import torch
from torch import nn
from torch_geometric.transforms import Compose, NormalizeScale, RandomFlip
from ml_lib.modules.geometric_blocks import SAModule, GlobalSAModule, MLP, FPModule
from ml_lib.modules.util import LightningBaseModule, F_x
+2 -4
View File
@@ -8,7 +8,6 @@ from datasets.shapenet import ShapeNetPartSegDataset
from models._point_net_2 import _PointNetCore
from utils.module_mixins import BaseValMixin, BaseTrainMixin, BaseOptimizerMixin, BaseDataloadersMixin, DatasetMixin
from utils.project_settings import GlobalVar
class PointNet2(BaseValMixin,
@@ -33,7 +32,7 @@ class PointNet2(BaseValMixin,
# This is not available with 6-dim cords
# RandomRotate(rot_max_angle, 0), RandomRotate(rot_max_angle, 1), RandomRotate(rot_max_angle, 2),
RandomTranslate(trans_max_distance),
NormalizeScale()
# NormalizeScale()
# NormalizePositions()
]
)
@@ -41,7 +40,6 @@ class PointNet2(BaseValMixin,
# Dataset
# =============================================================================
self.dataset = self.build_dataset(ShapeNetPartSegDataset,
collate_per_segment=True,
transform=transforms,
cluster_type=self.params.cluster_type,
refresh=self.params.refresh,
@@ -51,7 +49,7 @@ class PointNet2(BaseValMixin,
# Model Paramters
# =============================================================================
# Additional parameters
self.n_classes = len(GlobalVar.classes) if not self.params.poly_as_plane else (len(GlobalVar.classes) - 2)
self.n_classes = len(self.dataset.train_dataset.classes)
# Modules
self.lin3 = torch.nn.Linear(128, self.n_classes)