Browse Source

[MNT] unify index of ListData, remove seg of valid

pull/4/head
Gao Enhao 2 years ago
parent
commit
110b149455
6 changed files with 34 additions and 86 deletions
  1. +16
    -16
      abl/bridge/simple_bridge.py
  2. +1
    -1
      abl/evaluation/semantics_metric.py
  3. +2
    -2
      abl/evaluation/symbol_metric.py
  4. +2
    -2
      abl/learning/abl_model.py
  5. +1
    -1
      abl/utils/cache.py
  6. +12
    -64
      examples/mnist_add/mnist_add_example.ipynb

+ 16
- 16
abl/bridge/simple_bridge.py View File

@@ -25,7 +25,7 @@ class SimpleBridge(BaseBridge):

def predict(self, data_samples: ListData) -> Tuple[List[ndarray], List[ndarray]]:
self.model.predict(data_samples)
return data_samples["pred_idx"], data_samples.get("pred_prob", None)
return data_samples.pred_idx, data_samples.get("pred_prob", None)

def abduce_pseudo_label(
self,
@@ -34,7 +34,7 @@ class SimpleBridge(BaseBridge):
require_more_revision: int = 0,
) -> List[List[Any]]:
self.abducer.batch_abduce(data_samples, max_revision, require_more_revision)
return data_samples["abduced_pseudo_label"]
return data_samples.abduced_pseudo_label

def idx_to_pseudo_label(
self, data_samples: ListData, mapping: Optional[Dict] = None
@@ -45,7 +45,7 @@ class SimpleBridge(BaseBridge):
data_samples.pred_pseudo_label = [
[mapping[_idx] for _idx in sub_list] for sub_list in pred_idx
]
return data_samples["pred_pseudo_label"]
return data_samples.pred_pseudo_label

def pseudo_label_to_idx(
self, data_samples: ListData, mapping: Optional[Dict] = None
@@ -57,7 +57,7 @@ class SimpleBridge(BaseBridge):
for sub_list in data_samples.abduced_pseudo_label
]
data_samples.abduced_idx = abduced_idx
return data_samples["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()
@@ -104,16 +104,16 @@ class SimpleBridge(BaseBridge):

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")
self.model.save(save_path=osp.join(save_dir, f"model_checkpoint_loop_{loop + 1}.pth"))
self.model.save(
save_path=osp.join(save_dir, f"model_checkpoint_loop_{loop + 1}.pth")
)

def _valid(self, data_samples: ListData, batch_size: int = 128) -> None:
for seg_idx in range((len(data_samples) - 1) // batch_size + 1):
sub_data_samples = data_samples[seg_idx * batch_size : (seg_idx + 1) * batch_size]
self.predict(sub_data_samples)
self.idx_to_pseudo_label(sub_data_samples)
def _valid(self, data_samples: ListData) -> None:
self.predict(data_samples)
self.idx_to_pseudo_label(data_samples)

for metric in self.metric_list:
metric.process(sub_data_samples)
for metric in self.metric_list:
metric.process(data_samples)

res = dict()
for metric in self.metric_list:
@@ -123,12 +123,12 @@ class SimpleBridge(BaseBridge):
msg += k + f": {v:.3f} "
print_log(msg, logger="current")

def valid(self, valid_data: Union[ListData, DataSet], batch_size: int = 128) -> None:
def valid(self, valid_data: Union[ListData, DataSet]) -> None:
if not isinstance(valid_data, ListData):
data_samples = self.data_preprocess(*valid_data)
else:
data_samples = valid_data
self._valid(data_samples, batch_size=batch_size)
self._valid(data_samples)

def test(self, test_data: Union[ListData, DataSet], batch_size: int = 128) -> None:
self.valid(test_data, batch_size=batch_size)
def test(self, test_data: Union[ListData, DataSet]) -> None:
self.valid(test_data)

+ 1
- 1
abl/evaluation/semantics_metric.py View File

@@ -11,7 +11,7 @@ class SemanticsMetric(BaseMetric):

def process(self, data_samples: Sequence[dict]) -> None:
for data_sample in data_samples:
if self.kb.check_equal(data_sample, data_sample["Y"][0]):
if self.kb.check_equal(data_sample, data_sample.Y[0]):
self.results.append(1)
else:
self.results.append(0)


+ 2
- 2
abl/evaluation/symbol_metric.py View File

@@ -8,9 +8,9 @@ class SymbolMetric(BaseMetric):
super().__init__(prefix)

def process(self, data_samples: Sequence[dict]) -> None:
pred_pseudo_label = data_samples["pred_pseudo_label"]
pred_pseudo_label = data_samples.pred_pseudo_label

gt_pseudo_label = data_samples["gt_pseudo_label"]
gt_pseudo_label = data_samples.gt_pseudo_label

if not len(pred_pseudo_label) == len(gt_pseudo_label):
raise ValueError("lengthes of pred_pseudo_label and gt_pseudo_label should be equal")


+ 2
- 2
abl/learning/abl_model.py View File

@@ -69,11 +69,11 @@ class ABLModel:
if hasattr(model, "predict_proba"):
prob = model.predict_proba(X=data_X)
label = prob.argmax(axis=1)
prob = reform_idx(prob, data_samples["X"])
prob = reform_idx(prob, data_samples.X)
else:
prob = None
label = model.predict(X=data_X)
label = reform_idx(label, data_samples["X"])
label = reform_idx(label, data_samples.X)

data_samples.pred_idx = label
if prob is not None:


+ 1
- 1
abl/utils/cache.py View File

@@ -102,7 +102,7 @@ class Cache(Generic[K, T]):
log_dir = ABLLogger.get_current_instance().log_dir
cache_dir = osp.join(log_dir, "cache")
os.makedirs(cache_dir, exist_ok=True)
cache_path = osp.join(cache_dir, "cache.pth")
cache_path = osp.join(cache_dir, "abduce_by_search_cache_res.pth")
with open(cache_path, "wb") as file:
pickle.dump(self.cache_dict, file, protocol=pickle.HIGHEST_PROTOCOL)
print_log(f"Cache will be saved to {cache_path}", logger="current")


+ 12
- 64
examples/mnist_add/mnist_add_example.ipynb View File

@@ -2,7 +2,7 @@
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -24,17 +24,9 @@
},
{
"cell_type": "code",
"execution_count": 2,
"execution_count": null,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"11/15 21:35:55 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Abductive Learning on the MNIST Add example.\n"
]
}
],
"outputs": [],
"source": [
"# Initialize logger\n",
"print_log(\"Abductive Learning on the MNIST Add example.\", logger=\"current\")\n",
@@ -54,7 +46,7 @@
},
{
"cell_type": "code",
"execution_count": 3,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -79,7 +71,7 @@
},
{
"cell_type": "code",
"execution_count": 4,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -92,7 +84,7 @@
},
{
"cell_type": "code",
"execution_count": 5,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -118,7 +110,7 @@
},
{
"cell_type": "code",
"execution_count": 6,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -138,7 +130,7 @@
},
{
"cell_type": "code",
"execution_count": 7,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -156,7 +148,7 @@
},
{
"cell_type": "code",
"execution_count": 8,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -175,7 +167,7 @@
},
{
"cell_type": "code",
"execution_count": 9,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -192,53 +184,9 @@
},
{
"cell_type": "code",
"execution_count": 10,
"execution_count": null,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"11/15 21:36:18 - abl - \u001b[5m\u001b[4m\u001b[33mWARNING\u001b[0m - Transform used in the training phase will be used in prediction.\n",
"11/15 21:36:21 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [1/5] segment(train) [1/3] model loss is 1.80390\n",
"11/15 21:36:24 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [1/5] segment(train) [2/3] model loss is 1.41898\n",
"11/15 21:36:26 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [1/5] segment(train) [3/3] model loss is 1.08221\n",
"11/15 21:36:26 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation start: loop(val) [1]\n",
"11/15 21:36:27 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation ended, mnist_add/character_accuracy: 0.590 \n",
"11/15 21:36:27 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Saving model: loop(save) [1]\n",
"11/15 21:36:27 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Checkpoints will be saved to results/20231115_21_35_55/weights/model_checkpoint_loop_1.pth\n",
"11/15 21:36:29 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [2/5] segment(train) [1/3] model loss is 0.65210\n",
"11/15 21:36:31 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [2/5] segment(train) [2/3] model loss is 0.13546\n",
"11/15 21:36:32 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [2/5] segment(train) [3/3] model loss is 0.08060\n",
"11/15 21:36:32 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation start: loop(val) [2]\n",
"11/15 21:36:34 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation ended, mnist_add/character_accuracy: 0.982 \n",
"11/15 21:36:34 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Saving model: loop(save) [2]\n",
"11/15 21:36:34 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Checkpoints will be saved to results/20231115_21_35_55/weights/model_checkpoint_loop_2.pth\n",
"11/15 21:36:35 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [3/5] segment(train) [1/3] model loss is 0.06446\n",
"11/15 21:36:37 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [3/5] segment(train) [2/3] model loss is 0.05224\n",
"11/15 21:36:39 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [3/5] segment(train) [3/3] model loss is 0.05119\n",
"11/15 21:36:39 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation start: loop(val) [3]\n",
"11/15 21:36:40 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation ended, mnist_add/character_accuracy: 0.989 \n",
"11/15 21:36:40 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Saving model: loop(save) [3]\n",
"11/15 21:36:40 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Checkpoints will be saved to results/20231115_21_35_55/weights/model_checkpoint_loop_3.pth\n",
"11/15 21:36:42 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [4/5] segment(train) [1/3] model loss is 0.04667\n",
"11/15 21:36:44 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [4/5] segment(train) [2/3] model loss is 0.04027\n",
"11/15 21:36:45 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [4/5] segment(train) [3/3] model loss is 0.03672\n",
"11/15 21:36:45 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation start: loop(val) [4]\n",
"11/15 21:36:46 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation ended, mnist_add/character_accuracy: 0.990 \n",
"11/15 21:36:46 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Saving model: loop(save) [4]\n",
"11/15 21:36:46 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Checkpoints will be saved to results/20231115_21_35_55/weights/model_checkpoint_loop_4.pth\n",
"11/15 21:36:48 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [5/5] segment(train) [1/3] model loss is 0.03381\n",
"11/15 21:36:50 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [5/5] segment(train) [2/3] model loss is 0.03333\n",
"11/15 21:36:52 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - loop(train) [5/5] segment(train) [3/3] model loss is 0.03195\n",
"11/15 21:36:52 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation start: loop(val) [5]\n",
"11/15 21:36:53 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation ended, mnist_add/character_accuracy: 0.992 \n",
"11/15 21:36:53 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Saving model: loop(save) [5]\n",
"11/15 21:36:53 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Checkpoints will be saved to results/20231115_21_35_55/weights/model_checkpoint_loop_5.pth\n",
"11/15 21:36:53 - abl - \u001b[4m\u001b[37mINFO\u001b[0m - Evaluation ended, mnist_add/character_accuracy: 0.988 \n"
]
}
],
"outputs": [],
"source": [
"bridge.train(train_data, loops=5, segment_size=10000, save_interval=1, save_dir=weights_dir)\n",
"bridge.test(test_data)"


Loading…
Cancel
Save