| @@ -12,7 +12,7 @@ def get_data(file, get_pseudo_label): | |||||
| if get_pseudo_label: | if get_pseudo_label: | ||||
| Z = [] | Z = [] | ||||
| Y = [] | Y = [] | ||||
| img_dir = './datasets/hwf/data/Handwritten_Math_Symbols/' | |||||
| img_dir = './datasets/data/Handwritten_Math_Symbols/' | |||||
| with open(file) as f: | with open(file) as f: | ||||
| data = json.load(f) | data = json.load(f) | ||||
| for idx in range(len(data)): | for idx in range(len(data)): | ||||
| @@ -36,9 +36,9 @@ def get_data(file, get_pseudo_label): | |||||
| def get_hwf(train = True, get_pseudo_label = False): | def get_hwf(train = True, get_pseudo_label = False): | ||||
| if(train): | if(train): | ||||
| file = './datasets/hwf/data/expr_train.json' | |||||
| file = './datasets/data/expr_train.json' | |||||
| else: | else: | ||||
| file = './datasets/hwf/data/expr_test.json' | |||||
| file = './datasets/data/expr_test.json' | |||||
| return get_data(file, get_pseudo_label) | return get_data(file, get_pseudo_label) | ||||
| @@ -8,7 +8,7 @@ | |||||
| "source": [ | "source": [ | ||||
| "import sys\n", | "import sys\n", | ||||
| "\n", | "\n", | ||||
| "sys.path.append(\"../\")\n", | |||||
| "sys.path.append(\"../../\")\n", | |||||
| "\n", | "\n", | ||||
| "import torch.nn as nn\n", | "import torch.nn as nn\n", | ||||
| "import torch\n", | "import torch\n", | ||||
| @@ -21,8 +21,8 @@ | |||||
| "from abl.models.wabl_models import WABLBasicModel\n", | "from abl.models.wabl_models import WABLBasicModel\n", | ||||
| "\n", | "\n", | ||||
| "from models.nn import SymbolNet\n", | "from models.nn import SymbolNet\n", | ||||
| "from datasets.hwf.get_hwf import get_hwf\n", | |||||
| "from abl import framework_hed" | |||||
| "from datasets.get_hwf import get_hwf\n", | |||||
| "from abl import framework" | |||||
| ] | ] | ||||
| }, | }, | ||||
| { | { | ||||
| @@ -150,7 +150,7 @@ | |||||
| "outputs": [], | "outputs": [], | ||||
| "source": [ | "source": [ | ||||
| "# Train model\n", | "# Train model\n", | ||||
| "framework_hed.train(\n", | |||||
| "framework.train(\n", | |||||
| " model, abducer, train_data, test_data, loop_num=15, sample_num=5000, verbose=1\n", | " model, abducer, train_data, test_data, loop_num=15, sample_num=5000, verbose=1\n", | ||||
| ")\n", | ")\n", | ||||
| "\n", | "\n", | ||||
| @@ -175,7 +175,7 @@ | |||||
| "name": "python", | "name": "python", | ||||
| "nbconvert_exporter": "python", | "nbconvert_exporter": "python", | ||||
| "pygments_lexer": "ipython3", | "pygments_lexer": "ipython3", | ||||
| "version": "3.8.13" | |||||
| "version": "3.8.16" | |||||
| }, | }, | ||||
| "orig_nbformat": 4 | "orig_nbformat": 4 | ||||
| }, | }, | ||||
| @@ -1,6 +1,4 @@ | |||||
| import torch | |||||
| import torchvision | import torchvision | ||||
| from torch.utils.data import Dataset | |||||
| from torchvision.transforms import transforms | from torchvision.transforms import transforms | ||||
| def get_data(file, img_dataset, get_pseudo_label): | def get_data(file, img_dataset, get_pseudo_label): | ||||
| @@ -23,12 +21,12 @@ def get_data(file, img_dataset, get_pseudo_label): | |||||
| def get_mnist_add(train = True, get_pseudo_label = False): | def get_mnist_add(train = True, get_pseudo_label = False): | ||||
| transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081, ))]) | transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081, ))]) | ||||
| img_dataset = torchvision.datasets.MNIST(root='./datasets/mnist_add/', train=train, download=True, transform=transform) | |||||
| img_dataset = torchvision.datasets.MNIST(root='./datasets/', train=train, download=True, transform=transform) | |||||
| if train: | if train: | ||||
| file = './datasets/mnist_add/train_data.txt' | |||||
| file = './datasets/train_data.txt' | |||||
| else: | else: | ||||
| file = './datasets/mnist_add/test_data.txt' | |||||
| file = './datasets/test_data.txt' | |||||
| return get_data(file, img_dataset, get_pseudo_label) | return get_data(file, img_dataset, get_pseudo_label) | ||||