Browse Source

[MNT] resolve comments of code in simple_bridge.py

pull/1/head
Gao Enhao 2 years ago
parent
commit
2e97dfa8b1
1 changed files with 58 additions and 24 deletions
  1. +58
    -24
      abl/bridge/simple_bridge.py

+ 58
- 24
abl/bridge/simple_bridge.py View File

@@ -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)

Loading…
Cancel
Save