diff --git a/abl/bridge/simple_bridge.py b/abl/bridge/simple_bridge.py index 0a521b2..4e11369 100644 --- a/abl/bridge/simple_bridge.py +++ b/abl/bridge/simple_bridge.py @@ -44,18 +44,47 @@ class SimpleBridge(BaseBridge): data_samples.abduced_idx = abduced_idx return data_samples.abduced_idx - def data_preprocess(self, X: List[Any], gt_pseudo_label: List[Any], Y: List[Any]) -> ListData: - data_samples = ListData() - - data_samples.X = X - data_samples.gt_pseudo_label = gt_pseudo_label - data_samples.Y = Y + def data_preprocess( + self, + prefix: str, + data: Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]], + ) -> ListData: + if isinstance(data, ListData): + data_samples = data + if not ( + hasattr(data_samples, "X") + and hasattr(data_samples, "gt_pseudo_label") + and hasattr(data_samples, "Y") + ): + raise ValueError( + f"{prefix}data should have X, gt_pseudo_label and Y attribute but " + f"only {data_samples.all_keys()} are provided." + ) + else: + X, gt_pseudo_label, Y = data + data_samples = ListData(X=X, gt_pseudo_label=gt_pseudo_label, Y=Y) return data_samples + def concat_data_samples( + self, unlabel_data_samples: ListData, label_data_samples: Optional[ListData] + ) -> ListData: + if label_data_samples is None: + return unlabel_data_samples + + unlabel_data_samples.X = unlabel_data_samples.X + label_data_samples.X + unlabel_data_samples.abduced_pseudo_label = ( + unlabel_data_samples.abduced_pseudo_label + label_data_samples.gt_pseudo_label + ) + unlabel_data_samples.Y = unlabel_data_samples.Y + label_data_samples.Y + return unlabel_data_samples + def train( self, train_data: Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]], + label_data: Optional[ + Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]] + ] = None, val_data: Optional[ Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]] ] = None, @@ -65,19 +94,28 @@ class SimpleBridge(BaseBridge): save_interval: Optional[int] = None, save_dir: Optional[str] = None, ): - if isinstance(train_data, ListData): - data_samples = train_data + data_samples = self.data_preprocess("train", train_data) + + if label_data is not None: + label_data_samples = self.data_preprocess("label", label_data) else: - data_samples = self.data_preprocess(*train_data) + label_data_samples = None - if isinstance(segment_size, int) and segment_size == 0: - raise ValueError("segment_size should be positive.") + if val_data is not None: + val_data_samples = self.data_preprocess("val", val_data) + else: + val_data_samples = data_samples - if isinstance(segment_size, float): + if isinstance(segment_size, int): + if segment_size <= 0: + raise ValueError("segment_size should be positive.") + elif isinstance(segment_size, float): if 0 < segment_size <= 1: segment_size = int(segment_size * len(data_samples)) else: raise ValueError("segment_size should be in (0, 1].") + else: + raise ValueError("segment_size should be int or float.") for loop in range(loops): for seg_idx in range((len(data_samples) - 1) // segment_size + 1): @@ -86,7 +124,7 @@ class SimpleBridge(BaseBridge): f"[{(seg_idx + 1)}/{(len(data_samples) - 1) // segment_size + 1}] ", logger="current", ) - + sub_data_samples = data_samples[ seg_idx * segment_size : (seg_idx + 1) * segment_size ] @@ -94,15 +132,13 @@ class SimpleBridge(BaseBridge): self.idx_to_pseudo_label(sub_data_samples) self.abduce_pseudo_label(sub_data_samples) self.filter_pseudo_label(sub_data_samples) + self.concat_data_samples(sub_data_samples, label_data_samples) self.pseudo_label_to_idx(sub_data_samples) self.model.train(sub_data_samples) if (loop + 1) % eval_interval == 0 or loop == loops - 1: print_log(f"Evaluation start: loop(val) [{loop + 1}]", logger="current") - if val_data is not None: - self.valid(val_data) - else: - self.valid(train_data) + self._valid(val_data_samples) if save_interval is not None and ((loop + 1) % save_interval == 0 or loop == loops - 1): print_log(f"Saving model: loop(save) [{loop + 1}]", logger="current") @@ -127,16 +163,14 @@ class SimpleBridge(BaseBridge): def valid( self, - valid_data: Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]], + val_data: Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]], ) -> None: - if not isinstance(valid_data, ListData): - data_samples = self.data_preprocess(*valid_data) - else: - data_samples = valid_data - self._valid(data_samples) + val_data_samples = self.data_preprocess(val_data) + self._valid(val_data_samples) def test( self, test_data: Union[ListData, Tuple[List[List[Any]], Optional[List[List[Any]]], List[Any]]], ) -> None: - self.valid(test_data) + test_data_samples = self.data_preprocess("test", test_data) + self._valid(test_data_samples)