|
|
|
@@ -10,16 +10,14 @@ |
|
|
|
"import torch.nn as nn\n", |
|
|
|
"import torch\n", |
|
|
|
"\n", |
|
|
|
"from abl.reasoning.reasoner import ReasonerBase\n", |
|
|
|
"from abl.reasoning.kb import KBBase\n", |
|
|
|
"from abl.reasoning import ReasonerBase, KBBase\n", |
|
|
|
"from abl.learning import BasicNN, ABLModel\n", |
|
|
|
"from abl.bridge import SimpleBridge\n", |
|
|
|
"from abl.evaluation import SymbolMetric, ABLMetric\n", |
|
|
|
"from abl.utils import ABLLogger\n", |
|
|
|
"\n", |
|
|
|
"from abl.utils.plog import logger\n", |
|
|
|
"from abl.learning.basic_nn import BasicNN\n", |
|
|
|
"from abl.learning.abl_model import ABLModel\n", |
|
|
|
"\n", |
|
|
|
"from models.nn import SymbolNet\n", |
|
|
|
"from datasets.get_hwf import get_hwf\n", |
|
|
|
"from abl import framework" |
|
|
|
"from examples.models.nn import SymbolNet\n", |
|
|
|
"from datasets.get_hwf import get_hwf" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
@@ -29,7 +27,7 @@ |
|
|
|
"outputs": [], |
|
|
|
"source": [ |
|
|
|
"# Initialize logger\n", |
|
|
|
"recorder = logger()" |
|
|
|
"logger = ABLLogger.get_instance(\"abl\")" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
@@ -115,19 +113,30 @@ |
|
|
|
" optimizer,\n", |
|
|
|
" device,\n", |
|
|
|
" save_interval=1,\n", |
|
|
|
" save_dir=recorder.save_dir,\n", |
|
|
|
" save_dir=logger.save_dir,\n", |
|
|
|
" batch_size=128,\n", |
|
|
|
" num_epochs=10,\n", |
|
|
|
" recorder=recorder,\n", |
|
|
|
" num_epochs=3,\n", |
|
|
|
")" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
"cell_type": "code", |
|
|
|
"execution_count": null, |
|
|
|
"metadata": {}, |
|
|
|
"outputs": [], |
|
|
|
"source": [ |
|
|
|
"# Initialize ABL model\n", |
|
|
|
"# The main function of the ABL model is to serialize data and \n", |
|
|
|
"# provide a unified interface for different machine learning models\n", |
|
|
|
"model = ABLModel(base_model)" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
"attachments": {}, |
|
|
|
"cell_type": "markdown", |
|
|
|
"metadata": {}, |
|
|
|
"source": [ |
|
|
|
"### Use ABL model to join two parts" |
|
|
|
"### Metric" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
@@ -136,10 +145,8 @@ |
|
|
|
"metadata": {}, |
|
|
|
"outputs": [], |
|
|
|
"source": [ |
|
|
|
"# Initialize ABL model\n", |
|
|
|
"# The main function of the ABL model is to serialize data and \n", |
|
|
|
"# provide a unified interface for different machine learning models\n", |
|
|
|
"model = ABLModel(base_model)" |
|
|
|
"# Add metric\n", |
|
|
|
"metric = [SymbolMetric(prefix=\"hwf\"), ABLMetric(prefix=\"hwf\")]" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
@@ -166,7 +173,7 @@ |
|
|
|
"cell_type": "markdown", |
|
|
|
"metadata": {}, |
|
|
|
"source": [ |
|
|
|
"### Train and save" |
|
|
|
"### Bridge Machine Learning and Logic Reasoning" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
@@ -175,13 +182,25 @@ |
|
|
|
"metadata": {}, |
|
|
|
"outputs": [], |
|
|
|
"source": [ |
|
|
|
"# Train model\n", |
|
|
|
"framework.train(\n", |
|
|
|
" model, abducer, train_data, epochs=15, sample=500, verbose=1\n", |
|
|
|
")\n", |
|
|
|
"\n", |
|
|
|
"# Save results\n", |
|
|
|
"recorder.dump()" |
|
|
|
"bridge = SimpleBridge(model, abducer, metric)" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
"attachments": {}, |
|
|
|
"cell_type": "markdown", |
|
|
|
"metadata": {}, |
|
|
|
"source": [ |
|
|
|
"### Train and Test" |
|
|
|
] |
|
|
|
}, |
|
|
|
{ |
|
|
|
"cell_type": "code", |
|
|
|
"execution_count": null, |
|
|
|
"metadata": {}, |
|
|
|
"outputs": [], |
|
|
|
"source": [ |
|
|
|
"bridge.train(train_data, epochs=3, batch_size=1000)\n", |
|
|
|
"bridge.test(test_data)" |
|
|
|
] |
|
|
|
} |
|
|
|
], |
|
|
|
|