This commit is contained in:
1 parent
05756fcc7b
commit
c717b5e466
5 files changed
+1877
-1
No files matched your search
@@ -0,0 +1,290 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "8dbc4931-506b-4eb7-9049-11fda71fa2fd",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/q315433/micromamba/envs/pmf/lib/python3.12/site-packages/torch/__init__.py:749: UserWarning: torch.set_default_tensor_type() is deprecated as of PyTorch 2.1, please use torch.set_default_dtype() and torch.set_default_device() as alternatives. (Triggered internally at ../torch/csrc/tensor/python_tensor.cpp:431.)\n",
|
||||
" _C._set_default_tensor_type(t)\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import sys\n",
|
||||
"import torch\n",
|
||||
"from pyprojroot import here as project_root\n",
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"sys.path.insert(0, str(project_root()))\n",
|
||||
"\n",
|
||||
"from src.evaluation.utils import get_test_path, get_model\n",
|
||||
"from src.evaluation.eval import meta_test\n",
|
||||
"\n",
|
||||
"from src.train_utils.trainer import train_parser\n",
|
||||
"from src.models.feature_extractors.pretrained_fe import get_fe_metadata\n",
|
||||
"import torchvision.transforms as transforms\n",
|
||||
"from PIL import Image\n",
|
||||
"\n",
|
||||
"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "a90ff098-ad85-45db-8576-54ffe4c8a7cc",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def test_transform():\n",
|
||||
" def _convert_image_to_rgb(im):\n",
|
||||
" return im.convert('RGB')\n",
|
||||
"\n",
|
||||
" return transforms.Compose([\n",
|
||||
" #transforms.Resize(224),\n",
|
||||
" transforms.Resize(224),\n",
|
||||
" #transforms.CenterCrop(224),\n",
|
||||
" _convert_image_to_rgb,\n",
|
||||
" transforms.ToTensor(),\n",
|
||||
" transforms.Normalize(mean=torch.tensor([0.4815, 0.4578, 0.4082]), std=torch.tensor([0.2686, 0.2613, 0.2758])),\n",
|
||||
" ])\n",
|
||||
"\n",
|
||||
"preprocess = test_transform()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "3a400d41-cc0b-4af2-aafe-fb1c82bf21a2",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Defaulting to float32 dtype\n",
|
||||
"Loaded pretrained timm model vit_base_patch16_clip_224.openai\n",
|
||||
"../caml_pretrained_models/CAML_CLIP/model.pth\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import enum\n",
|
||||
"\n",
|
||||
"class T:\n",
|
||||
" fe_type = \"timm:vit_base_patch16_clip_224.openai:768\"\n",
|
||||
" #fe_type = \"timm:vit_huge_patch14_clip_224.laion2b:1280\"\n",
|
||||
" fe_dim = 768\n",
|
||||
" fe_dtype = \"float32\"\n",
|
||||
" model = \"CAML\"\n",
|
||||
" dropout = 0.0\n",
|
||||
" encoder_size = \"large\"\n",
|
||||
"\n",
|
||||
"fe_metadata = get_fe_metadata(T())\n",
|
||||
"#test_path = get_test_path(args, data_path)\n",
|
||||
"#device = torch.device(f'cuda:{args.gpu}')\n",
|
||||
"\n",
|
||||
"# Get the model and load its weights.\n",
|
||||
"model, model_path = get_model(T(), fe_metadata, device)\n",
|
||||
"print(model_path)\n",
|
||||
"#print(model)\n",
|
||||
"if model_path:\n",
|
||||
" model.load_state_dict(torch.load(model_path, map_location=f'cuda:0'), strict=False)\n",
|
||||
"model.to(device)\n",
|
||||
"_= model.eval()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"id": "534d1750-0433-403c-9d51-c0522747a97f",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"(0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2)\n",
|
||||
"torch.Size([65, 3, 224, 224])\n",
|
||||
"3\n",
|
||||
"62\n",
|
||||
"torch.Size([62, 4, 768])\n",
|
||||
"tensor([0, 1, 2], device='cuda:0')\n",
|
||||
"tensor([0, 1, 2], device='cuda:0')\n",
|
||||
"torch.Size([3])\n",
|
||||
"torch.Size([62, 4, 768])\n",
|
||||
"herre\n",
|
||||
"torch.Size([62])\n",
|
||||
"(62,)\n",
|
||||
"(0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2)\n",
|
||||
"torch.Size([65, 3, 224, 224])\n",
|
||||
"9\n",
|
||||
"56\n",
|
||||
"torch.Size([56, 10, 768])\n",
|
||||
"tensor([0, 0, 0, 1, 1, 1, 2, 2, 2], device='cuda:0')\n",
|
||||
"tensor([0, 0, 0, 1, 1, 1, 2, 2, 2], device='cuda:0')\n",
|
||||
"torch.Size([9])\n",
|
||||
"torch.Size([56, 10, 768])\n",
|
||||
"herre\n",
|
||||
"torch.Size([56])\n",
|
||||
"(56,)\n",
|
||||
"(0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2)\n",
|
||||
"torch.Size([65, 3, 224, 224])\n",
|
||||
"15\n",
|
||||
"50\n",
|
||||
"torch.Size([50, 16, 768])\n",
|
||||
"tensor([0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2], device='cuda:0')\n",
|
||||
"tensor([0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2], device='cuda:0')\n",
|
||||
"torch.Size([15])\n",
|
||||
"torch.Size([50, 16, 768])\n",
|
||||
"herre\n",
|
||||
"torch.Size([50])\n",
|
||||
"(50,)\n",
|
||||
"[0.58064516 0.51785714 0.52 ]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"\n",
|
||||
"img_path = \"../pmf_cvpr22/data_custom\"\n",
|
||||
"\n",
|
||||
"def filecnt_in_dir(dirr, typ):\n",
|
||||
" _, _, files = next(os.walk(f\"{img_path}/{dirr}/test/{typ}/\"))\n",
|
||||
" return len(files)\n",
|
||||
"\n",
|
||||
"def evaluate(shot, way, folder):\n",
|
||||
" ts = [\"good\", \"broken_small\", \"broken_large\", \"contamination\"]\n",
|
||||
" tss = [\"good\", \"cable_swap\", \"combined\", \"cut_inner_insulation\", \"cut_outer_insulation\", \"missing_cable\", \"missing_wire\", \"poke_insulation\"]\n",
|
||||
" tss = ts\n",
|
||||
" cat = [\"bottle\", \"cable\"]\n",
|
||||
"\n",
|
||||
" #goodnr = (len(tss)-1) * shot\n",
|
||||
" \n",
|
||||
" with torch.no_grad():\n",
|
||||
" #img_supp = [preprocess(Image.open(f\"{img_path}/{folder}/train/good/{i:03d}.png\")).unsqueeze(0).to(device) for i in range(shot)]\n",
|
||||
" img_supp = [preprocess(Image.open(f\"{img_path}/{folder}/test/{n}/{i:03d}.png\")).unsqueeze(0).to(device) for n in tss[1:4] for i in range(shot)]\n",
|
||||
" \n",
|
||||
" tmp = [(preprocess(Image.open(f\"{img_path}/{folder}/test/{n}/{i:03d}.png\")).unsqueeze(0).to(device), tss.index(n)-1) for n in tss[1:4] for i in range(shot, filecnt_in_dir(folder, n))]\n",
|
||||
" img_query, query_labels = zip(*tmp)\n",
|
||||
" #print(tmp)\n",
|
||||
" print(query_labels)\n",
|
||||
" \n",
|
||||
" img_concat = img_supp + list(img_query)\n",
|
||||
" img_concat = torch.cat(img_concat, 0)\n",
|
||||
" print(img_concat.shape)\n",
|
||||
" print(len(img_supp))\n",
|
||||
" print(len(img_query))\n",
|
||||
" #shot = (len(tss)-1) * shot\n",
|
||||
" \n",
|
||||
" #logits = model.meta_test(img_concat, way=4, shot=shot, query_shot=1)\n",
|
||||
" #print(logits)\n",
|
||||
" #\n",
|
||||
" feature_vector = model.get_feature_vector(img_concat)\n",
|
||||
" support_features = feature_vector[:way * shot]\n",
|
||||
" query_features = feature_vector[way * shot:]\n",
|
||||
" b, d = query_features.shape\n",
|
||||
" \n",
|
||||
" # Reshape query and support to a sequence.\n",
|
||||
" support = support_features.reshape(1, way * shot, d).repeat(b, 1, 1)\n",
|
||||
" query = query_features.reshape(-1, 1, d)\n",
|
||||
" feature_sequences = torch.cat([query, support], dim=1)\n",
|
||||
" print(feature_sequences.shape)\n",
|
||||
" \n",
|
||||
" #labels = torch.LongTensor([i // shot for i in range(shot * way)]).to(device)\n",
|
||||
" labels = torch.arange(way).repeat(shot, 1).T.flatten().to(model.device)\n",
|
||||
" print(labels)\n",
|
||||
" #labels = torch.from_numpy(np.ones(shape=shot, dtype=int)).to(device)\n",
|
||||
" #labels = torch.cat([torch.from_numpy(np.zeros(shape=shot, dtype=int)).to(device), labels])\n",
|
||||
" print(labels)\n",
|
||||
" \n",
|
||||
" #labels = torch.LongTensor([0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ]).to(device)\n",
|
||||
" print(labels.shape)\n",
|
||||
" print(feature_sequences.shape)\n",
|
||||
" logits = model.transformer_encoder.forward_imagenet_v2(feature_sequences, labels, way, shot)\n",
|
||||
" #print(logits)\n",
|
||||
" _, max_index = torch.max(logits[:, :way], 1)\n",
|
||||
" #print(max_index.cpu().numpy())\n",
|
||||
" #bbb = np.ones(shape=(14*4))\n",
|
||||
" #bbb[:14] = 0\n",
|
||||
" #print(np.mean(max_index.cpu().numpy() == bbb))\n",
|
||||
" print(\"herre\")\n",
|
||||
" print(max_index.shape)\n",
|
||||
" print(np.array(query_labels).shape)\n",
|
||||
"\n",
|
||||
" return np.mean(max_index.cpu().numpy() == np.array(query_labels))\n",
|
||||
"\n",
|
||||
"scores = [evaluate(shot, 3, \"bottle\") for shot in [1,3,5]]\n",
|
||||
"print(np.array(scores))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "22d68c6b-0df6-4953-8447-7843932fa974",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"CAML:\n",
|
||||
"Resulsts:\n",
|
||||
"\n",
|
||||
"bottle:\n",
|
||||
"jeweils 1,3,5 shots normal\n",
|
||||
"[0.40740741 0.39726027 0.30769231]\n",
|
||||
"\n",
|
||||
"inbalanced - mehr good shots 5,10,15,30 -> alle anderen nur 5\n",
|
||||
"- not possible\n",
|
||||
"1q\n",
|
||||
"2 ways nur detektieren ob fehlerhaft oder nicht 3,6,9 shots -> wegen model restrictions\n",
|
||||
"[0.79012346 0.84415584 0.87671233]\n",
|
||||
"\n",
|
||||
"inbalance 2 way 5,10,15,30 -> rest 5\n",
|
||||
"- not possible\n",
|
||||
"\n",
|
||||
"nur fehlerklasse erkennen 1,3,5\n",
|
||||
"[0.58064516 0.51785714 0.52 ]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"cable:\n",
|
||||
"jeweils 1,3,5 shots normal\n",
|
||||
"[0.24031008 0.19834711 0.15929204]\n",
|
||||
"\n",
|
||||
"inbalanced - mehr good shots 5,10,15,30 -> alle anderen nur 5\n",
|
||||
"- not possible\n",
|
||||
"\n",
|
||||
"2 ways nur detektieren ob fehlerhaft oder nicht 1,3,5 shots\n",
|
||||
"[0.57364341 0.54545455 0.59292035]\n",
|
||||
"\n",
|
||||
"inbalance 2 way 5,10,15,30 -> rest 5\n",
|
||||
"- not possible\n",
|
||||
"\n",
|
||||
"nur fehlerklasse erkennen 1,3,5\n",
|
||||
"[0.12962963 0.36363636 0.58823529]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.4"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABW4AAAPdCAYAAAAauvH/AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjguNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8fJSN1AAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3gU1frA8e/upvfeSKVJ74TelBJBFL0oigKhiIKK/LCAlSYiYAFUuBaqYNfrtSAYFEQuSGhBIKhgEkJJCKQS0ja78/tjkiWbTUgC6byf55mH3bNnZs4cUmbfvPsejaIoCkIIIYQQQgghhBBCCCHqDW1dD0AIIYQQQgghhBBCCCGEOQncCiGEEEIIIYQQQgghRD0jgVshhBBCCCGEEEIIIYSoZyRwK4QQQgghhBBCCCGEEPWMBG6FEEIIIYQQQgghhBCinpHArRBCCCGEEEIIIYQQQtQzErgVQgghhBBCCCGEEEKIekYCt0IIIYQQQgghhBBCCFHPSOBWCCGEEEIIIYQQQggh6hkJ3AohGrzLly/z7LPPMnToULy9vdFoNMybN6/S+8fExDBixAiCg4Oxt7fHw8ODXr16sWnTpkrtv3PnTjQaTZnb77//fp1XZS40NJQ77rijWo4FcP78eebNm0dMTEy1HbOxGjhwIAMHDqzrYQghhBA3tfXr16PRaEhISKjyvsX3al9++WW1jaf4mDt37qy2Y9aVefPmodFo6noYogLV+TU3cOBA2rVrd+ODKpKTk8O8efMaxfeDEPWNVV0PQAghblRqairvv/8+HTt2ZNSoUXz44YdV2j8jI4OgoCAeeOABmjRpwpUrV9i8eTPjxo0jISGBF198sVLHefXVVxk0aJBZW3XeEFWn8+fPM3/+fEJDQ+nUqVNdD0cIIYQQ4ppGjBjB3r178ff3r+uhCCFKycnJYf78+QCS8CBENZPArRCiwQsJCSE9PR2NRsOlS5eqHLgtK6PyjjvuID4+nvfff7/SgdsWLVrQs2fPKp1bVI2iKOTl5WFvb1/XQxFCCCFELfL29sbb27uuhyEaALlfFEI0JlIqQQjR4BWXJahuXl5eWFnV/N+34uLiuP/++wkICMDW1hZfX19uu+22MssYbN26lS5dumBvb0+rVq1Yu3atRZ9jx45x11134e7ujp2dHZ06dWLDhg2m13fu3En37t0BmDhxomn+istLVGU8JUVGRuLk5MSpU6cYPnw4Tk5OBAUF8dRTT5Gfn2/WNy0tjenTp9OkSRNsbGxo2rQpL7zwgkU/jUbD448/zr///W9at26Nra0tGzZsMH1c8pdffuHhhx/G09MTFxcXxo8fz5UrV0hOTua+++7Dzc0Nf39/nn76afR6vdmx58+fT48ePfDw8MDFxYUuXbqwZs0aFEW55nUKIYQQovaVVSqh+OPe+/fvp1+/fjg4ONC0aVNee+01jEajxTHy8vKYNWsWfn5+2NvbM2DAAA4fPmzW58CBA9x///2EhoZib29PaGgoDzzwAKdPn65wjJXdt/haduzYwbRp0/Dy8sLT05N77rmH8+fPWxz3448/plevXjg5OeHk5ESnTp1Ys2aNWZ/t27dz22234eLigoODA3369OHnn3+2ONYPP/xAp06dsLW1JSwsjNdff73C67oe1XGvdqP3iwAnT55k7Nix+Pj4YGtrS+vWrXn33XcrdQ0ZGRlMnjwZDw8PnJycGDFiBHFxcWWWZdu9eze33XYbzs7OODg40Lt3b3744QeLY1Z0n17szz//JCIiAgcHB7y8vHj00Ue5fPlypcZ98eJFpk6dSlBQELa2tnh7e9OnTx+2b99u0bcy3zuJiYk89NBDZnP4xhtvmPolJCSY/qgyf/5803uLyMjIKo9HCGFJMm6FEKKI0WjEaDSSnp7OF198wbZt23jnnXcqvf9jjz3G/fffj4ODA7169eKll16ib9++Fe43fPhwDAYDS5cuJTg4mEuXLrFnzx4yMjLM+h05coSnnnqKOXPm4Ovry4cffsjkyZNp3rw5/fv3B+Cvv/6id+/e+Pj4sHLlSjw9Pdm0aRORkZFcuHCBZ599li5durBu3TomTpzIiy++yIgRIwAIDAys0njKotfrufPOO5k8eTJPPfUUu3btYuHChbi6uvLyyy8D6pumQYMG8c8//zB//nw6dOjAb7/9xuLFi4mJibG4yf3mm2/47bffePnll/Hz88PHx4f9+/cDMGXKFO655x4+/fRTDh8+zPPPP09hYSF//fUX99xzD1OnTmX79u0sWbKEgIAAZs2aZTpuQkICjzzyCMHBwQD8/vvvPPHEE5w7d840ViGEEELUb8nJyTz44IM89dRTzJ07l//85z8899xzBAQEMH78eLO+zz//PF26dOHDDz8kMzOTefPmMXDgQA4fPkzTpk0B9f7glltu4f7778fDw4OkpCRWr15N9+7diY2NxcvLq9yxVHXfKVOmMGLECD7++GPOnDnDM888w0MPPcQvv/xi6vPyyy+zcOFC7rnnHp566ilcXV05duyYWTB406ZNjB8/nrvuuosNGzZgbW3Ne++9x7Bhw9i2bRu33XYbAD///DN33XUXvXr14tNPPzXd7124cOGG/x/Kc733atVxvxgbG0vv3r0JDg7mjTfewM/Pj23btjFjxgwuXbrE3Llzyx230Whk5MiRHDhwgHnz5tGlSxf27t1LRESERd9ff/2VIUOG0KFDB9asWYOtrS2rVq1i5MiRfPLJJ4wZMwao3H06wIULFxgwYADW1tasWrUKX19fNm/ezOOPP16pOR83bhyHDh1i0aJFtGzZkoyMDA4dOkRqaqpZv8p871y8eJHevXtTUFDAwoULCQ0N5fvvv+fpp5/mn3/+YdWqVfj7+7N161YiIiKYPHkyU6ZMATAFcys7HiFEORQhhGhELl68qADK3Llzq7zvI488ogAKoNjY2CirVq2q1H6HDh1SnnzySeU///mPsmvXLmXt2rVK69atFZ1Op2zduvWa+166dEkBlOXLl1+zX0hIiGJnZ6ecPn3a1Jabm6t4eHgojzzyiKnt/vvvV2xtbZXExESz/W+//XbFwcFBycjIUBRFUfbv368Ayrp1665rPGWZMGGCAiiff/65Wfvw4cOVW265xfT83//+d5n9lixZogDKTz/9ZGoDFFdXVyUtLc2s77p16xRAeeKJJ8zaR40apQDKm2++adbeqVMnpUuXLuWO3WAwKHq9XlmwYIHi6empGI1G02sDBgxQBgwYcO2LF0IIIUSNKv7dHx8fb2obMGCAAij79u0z69umTRtl2LBhpuc7duxQAKVLly5mv+MTEhIUa2trZcqUKeWet7CwUMnOzlYcHR2VFStWWBxzx44dVd63+FqmT59u1n/p0qUKoCQlJSmKoihxcXGKTqdTHnzwwXLPceXKFcXDw0MZOXKkWbvBYFA6duyohIeHm9p69OihBAQEKLm5uaa2rKwsxcPDQ6nu0MCN3qtVx/3isGHDlMDAQCUzM9Os/fHHH1fs7Ows+pf0ww8/KICyevVqs/bFixdbvNfo2bOn4uPjo1y+fNnUVlhYqLRr104JDAw0fc1V9j599uzZikajUWJiYsz6DRkypMKvOUVRFCcnJ2XmzJnX7FPZ7505c+aU2W/atGmKRqNR/vrrL0VRrv0erDLjEUKUT0olCCFEkeeff579+/fzww8/MGnSJB5//PFKfXysc+fOLF++nFGjRtGvXz8mTpzInj178Pf3N/3lvDweHh40a9aMZcuW8eabb3L48OEyP9oH0KlTJ1N2KICdnR0tW7Y0y7j45ZdfuO222wgKCjLbNzIykpycHPbu3Vtt4ymLRqNh5MiRZm0dOnSwGKOjoyOjR4+2GCNg8bG+W2+9FXd39zLPd8cdd5g9b926NYApi7hke+mPKf7yyy8MHjwYV1dXdDod1tbWvPzyy6SmppKSklLBlQohhBCiPvDz8yM8PNysrfS9R7GxY8ealdcKCQmhd+/e7Nixw9SWnZ3N7Nmzad68OVZWVlhZWeHk5MSVK1c4ceLENcdS1X3vvPNOi3EDprFHRUVhMBh47LHHyj3nnj17SEtLY8KECRQWFpo2o9FIREQE+/fv58qVK1y5coX9+/dzzz33YGdnZ9rf2dnZ4t6tLIqimB2/sLCwwn3g+u/VbvR+MS8vj59//pm7774bBwcHs3EPHz6cvLw8fv/993LH/euvvwJw3333mbU/8MADZs+vXLnCvn37GD16NE5OTqZ2nU7HuHHjOHv2LH/99Zfpmipzn75jxw7atm1Lx44dzfqNHTu23PGWFB4ezvr163nllVf4/fffLUpQFKvM984vv/xCmzZtLPpFRkaiKIpZdviNjkcIUTYJ3AohRJHg4GC6devG8OHDWb16NVOnTuW5557j4sWLVT6Wm5sbd9xxB3/88Qe5ubnl9tNoNPz8888MGzaMpUuX0qVLF7y9vZkxY4ZFHStPT0+L/W1tbc2On5qaWuZqywEBAabXr6Uq4ymLg4OD2ZuB4jHm5eWZjdHPz8+iLrGPjw9WVlYWY7zW6tEeHh5mz21sbMptLzmG6Ohohg4dCsAHH3zA//73P/bv388LL7wAcM3/MyGEEELUH5W5Pyrm5+dXZlvJe4+xY8fyzjvvMGXKFLZt20Z0dDT79+/H29u7wvuDqu5beuy2trbA1fuQ4nvQ4nJWZSkuczB69Gisra3NtiVLlqAoCmlpaaSnp2M0Gsudg4oUl2AouVXG9d6r3ej9YmpqKoWFhbz99tsW4x4+fDgAly5dKnfcqampWFlZWYzT19fX7Hl6ejqKolTq/ruy9+nF115aZf6fAD777DMmTJjAhx9+SK9evfDw8GD8+PEkJyeb9auN9xZVGY8QomxS41YIIcoRHh7Ov//9b+Li4q5rFWOlaJGrihZOCwkJMS0w8ffff/P5558zb948CgoK+Pe//12lc3p6epKUlGTRXrzQxbXqstXEeMob475Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1400x1000 with 6 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"\n",
|
||||
"# Define data\n",
|
||||
"bottle_data = {\n",
|
||||
" \"ResNet50\": {\n",
|
||||
" \"1,3,5 shots normal\": [0.5892857142857143, 0.7321428571428571, 0.75],\n",
|
||||
" \"inbalanced - more good shots\": [0.75, 0.732, 0.696, 0.696],\n",
|
||||
" \"2 ways only detect if faulty or not\": [0.8395, 0.8315, 0.8031],\n",
|
||||
" \"inbalance 2 way\": [0.8031, 0.81893, 0.8336, 0.8031],\n",
|
||||
" \"only faulty class detect\": [0.7638, 0.7428, 0.787]\n",
|
||||
" },\n",
|
||||
" \"P>M>F\": {\n",
|
||||
" \"1,3,5 shots normal\": [0.67910401, 0.71710526, 0.78860294],\n",
|
||||
" \"inbalanced - more good shots\": [0.78768382, 0.78860294, 0.75827206, 0.74356618],\n",
|
||||
" \"2 ways only detect if faulty or not\": [0.86422306, 0.93201754, 0.93933824],\n",
|
||||
" \"inbalance 2 way\": [0.92371324, 0.87867647, 0.86397059, 0.87775735],\n",
|
||||
" \"only faulty class detect\": [0.57380952, 0.76705653, 0.84191176]\n",
|
||||
" },\n",
|
||||
" \"CAML\": {\n",
|
||||
" \"1,3,5 shots normal\": [0.40740741, 0.39726027, 0.30769231],\n",
|
||||
" \"2 ways only detect if faulty or not\": [0.79012346, 0.84415584, 0.87671233],\n",
|
||||
" \"only faulty class detect\": [0.58064516, 0.51785714, 0.52]\n",
|
||||
" }\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"cable_data = {\n",
|
||||
" \"ResNet50\": {\n",
|
||||
" \"1,3,5 shots normal\": [0.21808, 0.43815, 0.4321478],\n",
|
||||
" \"inbalanced - more good shots\": [0.4321478, 0.432986, 0.42340, 0.4464635],\n",
|
||||
" \"2 ways only detect if faulty or not\": [0.8592, 0.8772, 0.8495],\n",
|
||||
" \"inbalance 2 way\": [0.8495, 0.8180, 0.7460, 0.6846],\n",
|
||||
" \"only faulty class detect\": [0.240, 0.4740, 0.4805]\n",
|
||||
" },\n",
|
||||
" \"P>M>F\": {\n",
|
||||
" \"1,3,5 shots normal\": [0.25199021, 0.44388328, 0.46975059],\n",
|
||||
" \"inbalanced - more good shots\": [0.50425859, 0.48023277, 0.43118282, 0.41842534],\n",
|
||||
" \"2 ways only detect if faulty or not\": [0.79263485, 0.8707712, 0.86756514],\n",
|
||||
" \"inbalance 2 way\": [0.86966158, 0.80142425, 0.80961366, 0.66028834],\n",
|
||||
" \"only faulty class detect\": [0.24383256, 0.43800505, 0.51304563]\n",
|
||||
" },\n",
|
||||
" \"CAML\": {\n",
|
||||
" \"1,3,5 shots normal\": [0.24031008, 0.19834711, 0.15929204],\n",
|
||||
" \"2 ways only detect if faulty or not\": [0.57364341, 0.54545455, 0.59292035],\n",
|
||||
" \"only faulty class detect\": [0.12962963, 0.36363636, 0.58823529]\n",
|
||||
" }\n",
|
||||
"}\n",
|
||||
"# Prepare the data\n",
|
||||
"measurement_types = [\n",
|
||||
" \"1,3,5 shots normal\",\n",
|
||||
" \"inbalanced - more good shots\",\n",
|
||||
" \"2 ways only detect if faulty or not\",\n",
|
||||
" \"inbalance 2 way\",\n",
|
||||
" \"only faulty class detect\"\n",
|
||||
"]\n",
|
||||
"\n",
|
||||
"models = [\"ResNet50\", \"P>M>F\", \"CAML\"]\n",
|
||||
"\n",
|
||||
"# Create subplots\n",
|
||||
"fig, axes = plt.subplots(3, 2, figsize=(14, 10))\n",
|
||||
"axes = axes.flatten()\n",
|
||||
"\n",
|
||||
"# Loop through each measurement type\n",
|
||||
"for i, measurement in enumerate(measurement_types):\n",
|
||||
" ax = axes[i]\n",
|
||||
" for model in models:\n",
|
||||
" # Get the bottle and cable data for the current measurement and model\n",
|
||||
" bottle_accuracy = bottle_data[model].get(measurement, [])\n",
|
||||
" cable_accuracy = cable_data[model].get(measurement, [])\n",
|
||||
" \n",
|
||||
" # Plot both bottle and cable data\n",
|
||||
" ax.plot(bottle_accuracy, marker='o', label=f'{model} - Bottle', linestyle='-')\n",
|
||||
" \n",
|
||||
" ax.set_title(measurement)\n",
|
||||
" ax.set_xlabel(\"Shots / Samples\")\n",
|
||||
" ax.set_ylabel(\"Accuracy\")\n",
|
||||
" ax.legend()\n",
|
||||
" ax.grid(True)\n",
|
||||
"\n",
|
||||
"# Adjust layout\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"image/png": "iVBORw0KGgoAAAANSUhEUgAABW0AAAPdCAYAAADxjUr8AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjguNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8fJSN1AAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd3gUVd/H4c/uplcICUloSegl9Bp670h5UZoCAoIiIAJSFIQAgoBUJYhKEUFFRVTKAwQhdAQEFAi9l9AhBRKy2Z33jyEry24aqSS/22su2bMzZ86cbDaz3z1zRqMoioIQQgghhBBCCCGEEEKIHEGb3Q0QQgghhBBCCCGEEEII8R8JbYUQQgghhBBCCCGEECIHkdBWCCGEEEIIIYQQQgghchAJbYUQQgghhBBCCCGEECIHkdBWCCGEEEIIIYQQQgghchAJbYUQQgghhBBCCCGEECIHkdBWCCGEEEIIIYQQQgghchAJbYUQQgghhBBCCCGEECIHkdBWCCGEEEIIIYQQQgghchAJbYUQL4Xo6GhGjx5Ny5Yt8fLyQqPRMGnSpFRvf/ToUdq1a0exYsVwdHTEw8ODoKAgVq5cmartw8LC0Gg0Vpf9+/e/4FGZ8/f3p3379hlSF8CNGzeYNGkSR48ezbA6c6vGjRvTuHHj7G6GEEIIIZ6zfPlyNBoNly5dSvO2iedvv/zyS4a1J7HOsLCwDKszu0yaNAmNRpPdzRApyMjXXOPGjQkMDEx/o556/PgxkyZNyhW/D0LkRDbZ3QAhhEiNe/fu8dVXX1G5cmU6derEN998k6btHz58SNGiRenRoweFCxfm0aNHrFq1ijfeeINLly4xfvz4VNUzbdo0mjRpYlaWkSc+GenGjRsEBwfj7+9PlSpVsrs5QgghhBBp1q5dO/bt24evr292N0UI8ZzHjx8THBwMIAMghMgEEtoKIV4Kfn5+PHjwAI1Gw927d9Mc2lobSdm+fXsuXrzIV199lerQtlSpUtSpUydN+xZpoygKcXFxODo6ZndThBBCCJHNvLy88PLyyu5miJeAnEMKIXIbmR5BCPFSSJyKIKN5enpiY5P5319duHCB7t27U6hQIezt7fH29qZZs2ZWpy7YtGkT1apVw9HRkbJly7J06VKLdY4fP07Hjh3Jnz8/Dg4OVKlShW+//db0fFhYGDVr1gTgzTffNPVf4pQSaWnPs/r27YuLiwvnzp2jbdu2uLi4ULRoUUaOHMmTJ0/M1r1//z6DBw+mcOHC2NnZUbx4cT766COL9TQaDUOGDOHLL7+kXLly2Nvb8+2335ouh9y2bRtvvfUWBQoUwM3Njd69e/Po0SNu3rzJa6+9Rr58+fD19WXUqFHo9XqzuoODg6lduzYeHh64ublRrVo1lixZgqIoyR6nEEIIIXIGa9MjJF7iffDgQRo0aICTkxPFixfn008/xWg0WtQRFxfHiBEj8PHxwdHRkUaNGnHkyBGzdQ4dOkT37t3x9/fH0dERf39/evToweXLl1NsY2q3TTyW7du388477+Dp6UmBAgXo0qULN27csKj3+++/JygoCBcXF1xcXKhSpQpLliwxW2fr1q00a9YMNzc3nJycqFevHn/++adFXRs2bKBKlSrY29sTEBDAZ599luJxvYiMOH9L7zkkwNmzZ+nZsycFCxbE3t6ecuXKsXDhwlQdw8OHD+nfvz8eHh64uLjQrl07Lly4YHV6tt27d9OsWTNcXV1xcnKibt26bNiwwaLOlM7dE506dYrWrVvj5OSEp6cnb7/9NtHR0alq9507dxg4cCBFixbF3t4eLy8v6tWrx9atWy3WTc3vzpUrV3j99dfN+nD27Nmm9S5dumT6QiU4ONj0eaNv375pbo8QwjoZaSuEyFOMRiNGo5EHDx7w888/s3nzZr744otUb//uu+/SvXt3nJycCAoKYsKECdSvXz/F7dq2bYvBYGDmzJkUK1aMu3fvsnfvXh4+fGi23j///MPIkSMZO3Ys3t7efPPNN/Tv35+SJUvSsGFDAE6fPk3dunUpWLAgCxYsoECBAqxcuZK+ffty69YtRo8eTbVq1Vi2bBlvvvkm48ePp127dgAUKVIkTe2xRq/X88orr9C/f39GjhzJzp07mTJlCu7u7nz88ceA+uGoSZMmnD9/nuDgYCpVqsSuXbuYPn06R48etTiZ/e2339i1axcff/wxPj4+FCxYkIMHDwIwYMAAunTpwo8//siRI0f48MMPSUhI4PTp03Tp0oWBAweydetWZsyYQaFChRgxYoSp3kuXLjFo0CCKFSsGwP79+xk6dCjXr183tVUIIYQQL5+bN2/Sq1cvRo4cycSJE1m7di3jxo2jUKFC9O7d22zdDz/8kGrVqvHNN98QGRnJpEmTaNy4MUeOHKF48eKAes5QpkwZunfvjoeHBxERESxatIiaNWsSHh6Op6dnkm1J67YDBgygXbt2fP/991y9epUPPviA119/nW3btpnW+fjjj5kyZQpdunRh5MiRuLu7c/z4cbMgeOXKlfTu3ZuOHTvy7bffYmtry+LFi2nVqhWbN2+mWbNmAPz555907NiRoKAgfvzxR9M54K1bt9L9c0jKi56/ZcQ5ZHh4OHXr1qVYsWLMnj0bHx8fNm/ezLBhw7h79y4TJ05Mst1Go5EOHTpw6NAhJk2aRLVq1di3bx+tW7e2WHfHjh20aNGCSpUqsWTJEuzt7QkJCaFDhw788MMPdOvWDUjduTvArVu3aNSoEba2toSEhODt7c2qVasYMmRIqvr8jTfe4PDhw3zyySeULl2ahw8fcvjwYe7du2e2Xmp+d+7cuUPdunWJj49nypQp+Pv7s379ekaNGsX58+cJCQnB19eXTZs20bp1a/r378+AAQMATEFuatsjhEiGIoQQL5k7d+4ogDJx4sQ0bzto0CAFUADFzs5OCQkJSdV2hw8fVt577z1l7dq1ys6dO5WlS5cq5cqVU3Q6nbJp06Zkt717964CKPPmzUt2PT8/P8XBwUG5fPmyqSw2Nlbx8PBQBg0aZCrr3r27Ym9vr1y5csVs+zZt2ihOTk7Kw4cPFUVRlIMHDyqAsmzZshdqjzV9+vRRAOWnn34yK2/btq1SpkwZ0+Mvv/zS6nozZsxQAGXLli2mMkBxd3dX7t+/b7busmXLFEAZOnSoWXmnTp0UQJkzZ45ZeZUqVZRq1aol2XaDwaDo9Xpl8uTJSoECBRSj0Wh6rlGjRkqjRo2SP3ghhBBCZLnE84GLFy+ayho1aqQAyl9//WW2bvny5ZVWrVqZHm/fvl0BlGrVqpn93b906ZJia2urDBgwIMn9JiQkKDExMYqzs7Myf/58izq3b9+e5m0Tj2Xw4MFm68+cOVMBlIiICEVRFOXChQuKTqdTevXqleQ+Hj16pHh4eCgdOnQwKzcYDErlypWVWrVqmcpq166tFCpUSImNjTWVRUVFKR4eHkpGRwLpPX/LiHPIVq1aKUWKFFEiIyPNyocMGaI4ODhYrP+sDRs2KICyaNEis/Lp06dbfP6oU6eOUrBgQSU6OtpUlpCQoAQGBipFihQxveZSe+4+ZswYRaPRKEePHjVbr0WLFim+5hRFUVxcXJThw4cnu05qf3fGjh1rdb133nlH0Wg0yunTpxVFSf5zWWraI4RInkyPIITIUz788EMOHjzIhg0b6NevH0OGDEnV5WFVq1Zl3rx5dOrUiQYNGvDmm2+yd+9efH19Td+OJ8XDw4MSJUowa9Ys5syZw5EjR6xeugdQpUoV06hQAAcHB0qXLm02qmLbtm00a9aMokWLmm3bt29fHj9+zL59+zKsPdZoNBo6dOhgVlapUiWLNjo7O9O1a1eLNgIWl+01bdqU/PnzW91f+/btzR6XK1cOwDR6+Nny5y9D3LZtG82bN8fd3R2dToetrS0ff/wx9+7d4/bt2ykcqRBCCCFyKh8fH2rVqmVW9vz5SKKePXuaTbPl5+dH3bp12b59u6ksJiaGMWPGULJkSWxsbLCxscHFxYVHjx5x8uTJZNuS1m1feeUVi3YDpraHhoZiMBh49913k9zn3r17uX//Pn369CEhIcG0GI1GWrduzcGDB3n06BGPHj3i4MGDdOnSBQcHB9P2rq6uFudz1iiKYlZ/QkJCitvAi5+/pfccMi4ujj///JPOnTvj5ORk1u62bdsSFxfH/v37k2z3jh07AHjttdfMynv06GH2+NGjR/z111907doVFxcXU7lOp+ONN97g2rVrnD592nRMqTl33759OxUqVKBy5cpm6/Xs2TPJ9j6rVq1aLF++nKlTp7J//36LaScSpeZ3Z9u2bZQvX95ivb59+6Ioitmo8PS2RwiRNAlthRB5SrFixahRowZt27Zl0aJFDBw4kHHjxnHnzp0015UvXz7at2/Pv//+S2xsbJLraTQa/vzzT1q1asXMmTOpVq0aXl5eDBs2zGKOqgIFClhsb29vb1b/vXv3rN5BuVChQqbnk5OW9ljj5ORkdtKf2Ma4uDizNvr4+FjMQ1ywYEFsbGws2pjcHaE9PDzMHtvZ2SVZ/mwbDhw4QMuWLQH4+uuv2bNnDwcPHuSjjz4CSPZnJoQQQoicLTXnTIl8fHyslj17PtKzZ0+++OILBgwYwObNmzlw4AAHDx7Ey8srxXOGtG77fNvt7e2B/85NEs9LE6e1siZxaoOuXbtia2trtsyYMQNFUbh//z4PHjzAaDQm2QcpSZx24dklNV70/C2955D37t0jISGBzz//3KLdbdu2BeDu3btJtvvevXvY2NhYtNPb29vs8YMHD1AUJVXn5Kk9d0889uel5ucEsHr1avr06cM333xDUFAQHh4e9O7dm5s3b5qtlxWfN9LSHiFE0mROWyFEnlarVi2+/PJLLly48EJ3Jlae3tAqpZuk+fn5mW4ccebMGX766ScmTZpEfHw8X375ZZr2WaBAASIiIizKE29gkdyca5nRnqTa+Ndff6Eoilnf3L5Line truncated
|
||||
"text/plain": [
|
||||
"<Figure size 1400x1000 with 6 Axes>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# Create subplots\n",
|
||||
"fig, axes = plt.subplots(3, 2, figsize=(14, 10))\n",
|
||||
"axes = axes.flatten()\n",
|
||||
"\n",
|
||||
"# Loop through each measurement type\n",
|
||||
"for i, measurement in enumerate(measurement_types):\n",
|
||||
" ax = axes[i]\n",
|
||||
" for model in models:\n",
|
||||
" # Get the bottle and cable data for the current measurement and model\n",
|
||||
" bottle_accuracy = bottle_data[model].get(measurement, [])\n",
|
||||
" cable_accuracy = cable_data[model].get(measurement, [])\n",
|
||||
" \n",
|
||||
" # Plot both bottle and cable data\n",
|
||||
" ax.plot(cable_accuracy, marker='o', label=f'{model} - Cable', linestyle='-')\n",
|
||||
" \n",
|
||||
" ax.set_title(measurement)\n",
|
||||
" ax.set_xlabel(\"Shots / Samples\")\n",
|
||||
" ax.set_ylabel(\"Accuracy\")\n",
|
||||
" ax.legend()\n",
|
||||
" ax.grid(True)\n",
|
||||
"\n",
|
||||
"# Adjust layout\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.14"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
@@ -0,0 +1,895 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import numpy as np\n",
|
||||
"import time\n",
|
||||
"import random\n",
|
||||
"import torch\n",
|
||||
"import torchvision.transforms as transforms\n",
|
||||
"#import gradio as gr\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"\n",
|
||||
"from models import get_model\n",
|
||||
"from dotmap import DotMap\n",
|
||||
"from PIL import Image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Pretrained weights found at dino_vitbase16_pretrain/dino_vitbase16_pretrain.pth\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"args = DotMap()\n",
|
||||
"args.deploy = 'finetune'\n",
|
||||
"args.arch = 'dino_base_patch16'\n",
|
||||
"args.no_pretrain = True\n",
|
||||
"small = \"https://huggingface.co/hushell/pmf_metadataset_dino/resolve/main/md_full_128x128_dinosmall_fp16_lr5e-5/best.pth?download=true\"\n",
|
||||
"full = 'https://huggingface.co/hushell/pmf_metadataset_dino/resolve/main/md_full_128x128_dinobase_fp16_lr5e-5/best.pth?download=true'\n",
|
||||
"args.resume = full\n",
|
||||
"args.api_key = 'AIzaSyAFkOGnXhy-2ZB0imDvNNqf2rHb98vR_qY'\n",
|
||||
"args.cx = '06d75168141bc47f1'\n",
|
||||
"\n",
|
||||
"args.ada_steps = 100\n",
|
||||
"#args.ada_lr= 0.0001\n",
|
||||
"#args.aug_prob = .95\n",
|
||||
"args.ada_lr= 0.0001\n",
|
||||
"args.aug_prob = .9\n",
|
||||
"args.aug_types = [\"color\", \"translation\"]\n",
|
||||
"\n",
|
||||
"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
|
||||
"model = get_model(args)\n",
|
||||
"model.to(device)\n",
|
||||
"checkpoint = torch.hub.load_state_dict_from_url(args.resume, map_location='cpu')\n",
|
||||
"model.load_state_dict(checkpoint['model'], strict=True)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# image transforms\n",
|
||||
"def test_transform():\n",
|
||||
" def _convert_image_to_rgb(im):\n",
|
||||
" return im.convert('RGB')\n",
|
||||
"\n",
|
||||
" return transforms.Compose([\n",
|
||||
" transforms.Resize(224),\n",
|
||||
" #transforms.CenterCrop(224),\n",
|
||||
" _convert_image_to_rgb,\n",
|
||||
" transforms.ToTensor(),\n",
|
||||
" transforms.Normalize(mean=[0.485, 0.456, 0.406],\n",
|
||||
" std=[0.229, 0.224, 0.225]),\n",
|
||||
" ])\n",
|
||||
"\n",
|
||||
"preprocess = test_transform()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 42,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def build_sup_set(shotnr, type_name, binary, good_sample_nr):\n",
|
||||
" classes = next(os.walk(f'data_custom/{type_name}/test'))[1]\n",
|
||||
" classes.remove(\"good\")\n",
|
||||
"\n",
|
||||
" supp_x = []\n",
|
||||
" supp_y = []\n",
|
||||
" mapping = {\n",
|
||||
" \"good\" : 0\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
" # add good manually\n",
|
||||
" x_good = [Image.open(f\"data_custom/{type_name}/train/good/{x:03d}.png\") for x in range(0, good_sample_nr)]\n",
|
||||
" supp_x.extend([preprocess(x) for x in x_good]) # (3, H, W))\n",
|
||||
" supp_y.extend([0] * good_sample_nr)\n",
|
||||
" \n",
|
||||
" for i,c in enumerate(classes):\n",
|
||||
" #i-=1\n",
|
||||
" x_im = [Image.open(f\"data_custom/{type_name}/test/{c}/{x:03d}.png\") for x in range(0, shotnr)]\n",
|
||||
" supp_x.extend([preprocess(x) for x in x_im]) # (3, H, W))\n",
|
||||
" if binary:\n",
|
||||
" supp_y.extend([1] * shotnr)\n",
|
||||
" mapping[\"anomaly\"] = 1\n",
|
||||
" else:\n",
|
||||
" supp_y.extend([i+1] * shotnr)\n",
|
||||
" mapping[c] = i+1\n",
|
||||
" \n",
|
||||
" supp_x = torch.stack(supp_x, dim=0).unsqueeze(0).to(device) # (1, n_supp*n_labels, 3, H, W)\n",
|
||||
" supp_y = torch.tensor(supp_y).long().unsqueeze(0).to(device) # (1, n_supp*n_labels)\n",
|
||||
" return supp_x, supp_y, mapping\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"metadata": {
|
||||
"scrolled": true
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def build_test_set(shotnr, keyy, type):\n",
|
||||
" _, _, files = next(os.walk(f\"data_custom/cable/test/{type}/\"))\n",
|
||||
" file_count = len(files)\n",
|
||||
" print(file_count)\n",
|
||||
"\n",
|
||||
" queries = [preprocess(Image.open(f\"data_custom/cable/test/{type}/{i:03d}.png\")).unsqueeze(0).unsqueeze(0).to(device) for i in range(shotnr,file_count)]\n",
|
||||
" labels = [keyy for x in range(shotnr,file_count)]\n",
|
||||
" return queries, labels\n",
|
||||
"\n",
|
||||
"def test(type, keyy, shotnr, folder):\n",
|
||||
" predictions = []\n",
|
||||
" _, _, files = next(os.walk(f\"data_custom/{folder}/test/{type}/\"))\n",
|
||||
" file_count = len(files)\n",
|
||||
" print(file_count)\n",
|
||||
"\n",
|
||||
" queries = [preprocess(Image.open(f\"data_custom/{folder}/test/{type}/{i:03d}.png\")).unsqueeze(0).unsqueeze(0).to(device) for i in range(shotnr,file_count)]\n",
|
||||
" queries = torch.cat(queries)\n",
|
||||
" with torch.cuda.amp.autocast(True):\n",
|
||||
" output = model(supp_x, supp_y, queries) # (1, 1, n_labels)\n",
|
||||
"\n",
|
||||
" probs = output.softmax(dim=-1).detach().cpu().numpy()\n",
|
||||
" predictions = np.argmax(probs, axis=2)\n",
|
||||
" print()\n",
|
||||
" return np.mean([x == keyy for x in predictions])\n",
|
||||
" pass\n",
|
||||
" \n",
|
||||
"#def test2(folder):\n",
|
||||
"# accs = []\n",
|
||||
"# queries = []\n",
|
||||
"# labels = []\n",
|
||||
"# for t in next(os.walk(f'data_custom/cable/test'))[1]:\n",
|
||||
"# q, l = build_test_set(shots, types.get(t, 1), t)\n",
|
||||
"# queries+=q\n",
|
||||
"# labels+=l\n",
|
||||
"#\n",
|
||||
"# queries = torch.cat(queries)\n",
|
||||
"# labels = np.array(labels)\n",
|
||||
"#\n",
|
||||
"# with torch.cuda.amp.autocast(True):\n",
|
||||
"# output = model(supp_x, supp_y, queries) # (1, 1, n_labels)\n",
|
||||
"#\n",
|
||||
"# probs = output.softmax(dim=-1).detach().cpu().numpy()\n",
|
||||
"# predictions = np.argmax(probs, axis=2)\n",
|
||||
"# print()\n",
|
||||
"# return np.mean([predictions == labels])\n",
|
||||
"# pass\n",
|
||||
"\n",
|
||||
"#print(f\"overall accuracy: {test(\"cable\")}\")\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 48,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"14\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry9: loss = 0.1475423127412796: 100%|██| 100/100 [00:29<00:00, 3.40it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_inner_insulation = 1.0\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry5: loss = 0.20609889924526215: 100%|█| 100/100 [00:29<00:00, 3.37it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for poke_insulation = 1.0\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry7: loss = 0.12025140225887299: 100%|█| 100/100 [00:29<00:00, 3.34it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cable_swap = 0.8571428571428571\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry5: loss = 0.2130972295999527: 100%|██| 100/100 [00:30<00:00, 3.30it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_outer_insulation = 1.0\n",
|
||||
"58\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry53: loss = 0.13926956057548523: 100%|█| 100/100 [00:30<00:00, 3.30it/s\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for good = 0.16981132075471697\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry7: loss = 0.16337624192237854: 100%|█| 100/100 [00:30<00:00, 3.28it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_cable = 1.0\n",
|
||||
"11\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry6: loss = 0.16593313217163086: 100%|█| 100/100 [00:30<00:00, 3.27it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for combined = 1.0\n",
|
||||
"13\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry8: loss = 0.16560573875904083: 100%|█| 100/100 [00:30<00:00, 3.27it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for bent_wire = 1.0\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp45, nQry5: loss = 0.18611018359661102: 100%|█| 100/100 [00:30<00:00, 3.28it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_wire = 0.8\n",
|
||||
"overall accuracy: 0.8696615753219527\n",
|
||||
"14\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry9: loss = 0.3357824385166168: 100%|██| 100/100 [00:33<00:00, 2.96it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_inner_insulation = 0.7777777777777778\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry5: loss = 0.3290153741836548: 100%|██| 100/100 [00:33<00:00, 2.96it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for poke_insulation = 0.6\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry7: loss = 0.22177687287330627: 100%|█| 100/100 [00:33<00:00, 2.96it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cable_swap = 0.8571428571428571\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry5: loss = 0.299775630235672: 100%|███| 100/100 [00:33<00:00, 2.96it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_outer_insulation = 1.0\n",
|
||||
"58\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry53: loss = 0.31954386830329895: 100%|█| 100/100 [00:33<00:00, 2.98it/s\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for good = 0.32075471698113206\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry7: loss = 0.336273193359375: 100%|███| 100/100 [00:33<00:00, 2.98it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_cable = 0.8571428571428571\n",
|
||||
"11\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry6: loss = 0.3643767237663269: 100%|██| 100/100 [00:33<00:00, 2.98it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for combined = 1.0\n",
|
||||
"13\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry8: loss = 0.3085792660713196: 100%|██| 100/100 [00:33<00:00, 2.98it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for bent_wire = 1.0\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp50, nQry5: loss = 0.34715649485588074: 100%|█| 100/100 [00:33<00:00, 2.98it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_wire = 0.8\n",
|
||||
"overall accuracy: 0.8014242454494026\n",
|
||||
"14\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry9: loss = 0.375447154045105: 100%|███| 100/100 [00:36<00:00, 2.76it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_inner_insulation = 0.6666666666666666\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry5: loss = 0.42370423674583435: 100%|█| 100/100 [00:36<00:00, 2.75it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for poke_insulation = 1.0\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry7: loss = 0.3982161581516266: 100%|██| 100/100 [00:36<00:00, 2.74it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cable_swap = 0.8571428571428571\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry5: loss = 0.3903641104698181: 100%|██| 100/100 [00:36<00:00, 2.75it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_outer_insulation = 1.0\n",
|
||||
"58\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry53: loss = 0.4019339382648468: 100%|█| 100/100 [00:36<00:00, 2.75it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for good = 0.41509433962264153\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry7: loss = 0.4283098876476288: 100%|██| 100/100 [00:36<00:00, 2.75it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_cable = 0.7142857142857143\n",
|
||||
"11\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry6: loss = 0.3741377890110016: 100%|██| 100/100 [00:36<00:00, 2.74it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for combined = 0.8333333333333334\n",
|
||||
"13\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry8: loss = 0.3858358860015869: 100%|██| 100/100 [00:36<00:00, 2.75it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for bent_wire = 1.0\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp55, nQry5: loss = 0.3570959270000458: 100%|██| 100/100 [00:36<00:00, 2.74it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_wire = 0.8\n",
|
||||
"overall accuracy: 0.8096136567834681\n",
|
||||
"14\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry9: loss = 0.5021733045578003: 100%|██| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_inner_insulation = 0.5555555555555556\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry5: loss = 0.5203520059585571: 100%|██| 100/100 [00:45<00:00, 2.20it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for poke_insulation = 0.4\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry7: loss = 0.524366021156311: 100%|███| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cable_swap = 0.42857142857142855\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry5: loss = 0.5256413221359253: 100%|██| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for cut_outer_insulation = 1.0\n",
|
||||
"58\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry53: loss = 0.5186663866043091: 100%|█| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for good = 0.7358490566037735\n",
|
||||
"12\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry7: loss = 0.5123675465583801: 100%|██| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_cable = 0.7142857142857143\n",
|
||||
"11\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry6: loss = 0.5076506733894348: 100%|██| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for combined = 0.8333333333333334\n",
|
||||
"13\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry8: loss = 0.490247517824173: 100%|███| 100/100 [00:45<00:00, 2.21it/s]\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for bent_wire = 0.875\n",
|
||||
"10\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"lr0.0001, nSupp70, nQry5: loss = 0.3723257780075073: 100%|██| 100/100 [00:45<00:00, 2.21it/s]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n",
|
||||
"accuracy for missing_wire = 0.4\n",
|
||||
"overall accuracy: 0.6602883431499785\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"#bottle_accs = []\n",
|
||||
"cable_accs = []\n",
|
||||
"\n",
|
||||
"for nr in [5, 10, 15, 30]:\n",
|
||||
" folder = \"cable\"\n",
|
||||
" shot = 5\n",
|
||||
" supp_x, supp_y, types = build_sup_set(shot, folder, True, nr)\n",
|
||||
" accs = []\n",
|
||||
" for t in next(os.walk(f'data_custom/{folder}/test'))[1]:\n",
|
||||
" #if t == \"good\":\n",
|
||||
" # continue\n",
|
||||
" accuracy = test(t, types.get(t, 1), shot, folder)\n",
|
||||
" print(f\"accuracy for {t} = {accuracy}\")\n",
|
||||
" accs.append(accuracy)\n",
|
||||
" print(f\"overall accuracy: {np.mean(accs)}\")\n",
|
||||
" cable_accs.append(np.mean(accs))\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 39,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[0.57380952 0.76705653 0.84191176]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(np.array(bottle_accs))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 49,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"[0.86966158 0.80142425 0.80961366 0.66028834]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(np.array(cable_accs))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"P>M>F:\n",
|
||||
"Resulsts:\n",
|
||||
"\n",
|
||||
"bottle:\n",
|
||||
"jeweils 1,3,5 shots normal\n",
|
||||
"[0.67910401 0.71710526 0.78860294]\n",
|
||||
"\n",
|
||||
"inbalanced - mehr good shots 5,10,15,30 -> alle anderen nur 5\n",
|
||||
"[0.78768382 0.78860294 0.75827206 0.74356618]\n",
|
||||
"\n",
|
||||
"2 ways nur detektieren ob fehlerhaft oder nicht 1,3,5 shots\n",
|
||||
"[0.86422306 0.93201754 0.93933824]\n",
|
||||
"\n",
|
||||
"inbalance 2 way 5,10,15,30 -> rest 5\n",
|
||||
"[0.92371324 0.87867647 0.86397059 0.87775735]\n",
|
||||
"\n",
|
||||
"nur fehlerklasse erkennen 1,3,5\n",
|
||||
"[0.57380952 0.76705653 0.84191176]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"cable:\n",
|
||||
"jeweils 1,3,5 shots normal\n",
|
||||
"[0.25199021 0.44388328 0.46975059]\n",
|
||||
"\n",
|
||||
"inbalanced - mehr good shots 5,10,15,30 -> alle anderen nur 5\n",
|
||||
"[0.50425859 0.48023277 0.43118282 0.41842534]\n",
|
||||
"\n",
|
||||
"2 ways nur detektieren ob fehlerhaft oder nicht 1,3,5 shots\n",
|
||||
"[0.79263485 0.8707712 0.86756514]\n",
|
||||
"\n",
|
||||
"inbalance 2 way 5,10,15,30 -> rest 5\n",
|
||||
"[0.86966158 0.80142425 0.80961366 0.66028834]\n",
|
||||
"\n",
|
||||
"nur fehlerklasse erkennen 1,3,5\n",
|
||||
"[0.24383256 0.43800505 0.51304563]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.4"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 4
|
||||
}
|
||||
@@ -0,0 +1,495 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"imports imported\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"#import numpy as np\n",
|
||||
"import numpy as np\n",
|
||||
"import torch\n",
|
||||
"import cv2\n",
|
||||
"\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from torch import optim, nn\n",
|
||||
"import torchvision\n",
|
||||
"from torchvision import datasets, models, transforms\n",
|
||||
"import albumentations as A\n",
|
||||
"from albumentations.pytorch import ToTensorV2\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"print(\"imports imported\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class Identity(nn.Module):\n",
|
||||
" def __init__(self):\n",
|
||||
" super(Identity, self).__init__()\n",
|
||||
" \n",
|
||||
" def forward(self, x):\n",
|
||||
" return x"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"ResNet(\n",
|
||||
" (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)\n",
|
||||
" (layer1): Sequential(\n",
|
||||
" (0): Bottleneck(\n",
|
||||
" (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" (downsample): Sequential(\n",
|
||||
" (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (1): Bottleneck(\n",
|
||||
" (conv1): Conv2d(256, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (2): Bottleneck(\n",
|
||||
" (conv1): Conv2d(256, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (layer2): Sequential(\n",
|
||||
" (0): Bottleneck(\n",
|
||||
" (conv1): Conv2d(256, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" (downsample): Sequential(\n",
|
||||
" (0): Conv2d(256, 512, kernel_size=(1, 1), stride=(2, 2), bias=False)\n",
|
||||
" (1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (1): Bottleneck(\n",
|
||||
" (conv1): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (2): Bottleneck(\n",
|
||||
" (conv1): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (3): Bottleneck(\n",
|
||||
" (conv1): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (layer3): Sequential(\n",
|
||||
" (0): Bottleneck(\n",
|
||||
" (conv1): Conv2d(512, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" (downsample): Sequential(\n",
|
||||
" (0): Conv2d(512, 1024, kernel_size=(1, 1), stride=(2, 2), bias=False)\n",
|
||||
" (1): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (1): Bottleneck(\n",
|
||||
" (conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (2): Bottleneck(\n",
|
||||
" (conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (3): Bottleneck(\n",
|
||||
" (conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (4): Bottleneck(\n",
|
||||
" (conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (5): Bottleneck(\n",
|
||||
" (conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (layer4): Sequential(\n",
|
||||
" (0): Bottleneck(\n",
|
||||
" (conv1): Conv2d(1024, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" (downsample): Sequential(\n",
|
||||
" (0): Conv2d(1024, 2048, kernel_size=(1, 1), stride=(2, 2), bias=False)\n",
|
||||
" (1): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (1): Bottleneck(\n",
|
||||
" (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" (2): Bottleneck(\n",
|
||||
" (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n",
|
||||
" (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",
|
||||
" (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",
|
||||
" (relu): ReLU(inplace=True)\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
" (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))\n",
|
||||
" (fc): Linear(in_features=2048, out_features=1000, bias=True)\n",
|
||||
")\n",
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"resnet50 = models.resnet50(weights=models.ResNetshotnr0_Weights.DEFAULT)\n",
|
||||
"\n",
|
||||
"print(resnetshotnr0)\n",
|
||||
"# Step 2: Modify the model to output features from the layer before the fully connected layer\n",
|
||||
"class ResNetshotnr0Embeddings(nn.Module):\n",
|
||||
" def __init__(self, original_model, layernr):\n",
|
||||
" super(ResNetshotnr0Embeddings, self).__init__()\n",
|
||||
" #print(list(original_model.children())[4 + layernr])\n",
|
||||
" #print(nn.Sequential(*list(original_model.children())[:4 + shotnr]))\n",
|
||||
" self.features = nn.Sequential(*list(original_model.children())[:4+layernr])\n",
|
||||
" #self.features = nn.Sequential(*list(original_model.children())[:-1]) # Exclude the fully connected layer\n",
|
||||
"\n",
|
||||
" def forward(self, x):\n",
|
||||
" x = self.features(x)\n",
|
||||
" x = torch.flatten(x, 1) # Flatten the tensor to (batch_size, 2048)\n",
|
||||
" return x\n",
|
||||
"\n",
|
||||
"# Instantiate the modified model\n",
|
||||
"model = ResNetshotnr0Embeddings(resnetshotnr0, shotnr) # 3 = layer before fully connected one\n",
|
||||
"model.eval() # Set the model to evaluation mode\n",
|
||||
"print()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Test"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 69,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"...............\n",
|
||||
"accuracy for broken_large = 0.6666666666666666\n",
|
||||
".................\n",
|
||||
"accuracy for broken_small = 0.8823529411764706\n",
|
||||
"................\n",
|
||||
"accuracy for contamination = 0.8125\n",
|
||||
"overall accuracy: 0.7871732026143791\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from sklearn.metrics.pairwise import cosine_similarity,euclidean_distances\n",
|
||||
"from metric_learn import LMNN,NCA\n",
|
||||
"import math\n",
|
||||
"\n",
|
||||
"pipe = A.Compose([A.Resize(256,256), A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), ToTensorV2()])\n",
|
||||
"#pipe = A.Compose([A.Resize(256,256), ToTensorV2()])\n",
|
||||
"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n",
|
||||
"\n",
|
||||
"m = ResNet50Embeddings(resnet50, 5) # 5 = all without fully ocnnected\n",
|
||||
"m.eval()\n",
|
||||
"m.to(device)\n",
|
||||
"\n",
|
||||
"def read_img(path):\n",
|
||||
" img = cv2.imread(path, cv2.IMREAD_COLOR)\n",
|
||||
" img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n",
|
||||
" #plt.imshow(img)\n",
|
||||
"\n",
|
||||
" imgpiped = pipe(image=img)[\"image\"].unsqueeze(0)\n",
|
||||
" return imgpiped\n",
|
||||
"\n",
|
||||
"def compare_embeddings(emb1, emb2, distance_metric):\n",
|
||||
" #cosi = torch.nn.CosineSimilarity(dim=0) \n",
|
||||
" #output = cosine_similarity([emb1.flatten(), emb2.flatten()])\n",
|
||||
" #output = euclidean_distances([emb1.flatten(), emb2.flatten()], [emb1.flatten(), emb2.flatten()])\n",
|
||||
" output = distance_metric(emb1, emb2)\n",
|
||||
" return output\n",
|
||||
"\n",
|
||||
"def merge_embeddings(embeddings):\n",
|
||||
" # todo calc cluster center or similar\n",
|
||||
" return np.average(embeddings, axis=0)\n",
|
||||
"\n",
|
||||
"import os\n",
|
||||
"#embedding_good_1 = m(read_img(f\"./data/bottle/test/good/001.png\")).detach().numpy()\n",
|
||||
"#embedding_good_2 = m(read_img(f\"./data/bottle/test/good/002.png\")).detach().numpy()\n",
|
||||
"#embedding_good = merge_embeddings([embedding_good_1, embedding_good_2])\n",
|
||||
"#embedding_contermination_1 = m(read_img(f\"./data/bottle/test/contamination/001.png\")).detach().numpy()\n",
|
||||
"#embedding_contermination_2 = m(read_img(f\"./data/bottle/test/contamination/002.png\")).detach().numpy()\n",
|
||||
"#embedding_contermination = merge_embeddings([embedding_contermination_1, embedding_contermination_2])\n",
|
||||
"#embedding_broken_small_1 = m(read_img(f\"./data/bottle/test/broken_small/001.png\")).detach().numpy()\n",
|
||||
"\n",
|
||||
"#embeddings_test = m(read_img(f\"./data/bottle/test/contamination/004.png\")).detach().numpy()\n",
|
||||
"\n",
|
||||
"#score = compare_embeddings(embedding_good_1, embeddings_test)\n",
|
||||
"\n",
|
||||
"#def calc_base_emb(t):\n",
|
||||
"# base_emb_1 = m(read_img(f\"./data/bottle/test/{t}/000.png\")).detach().numpy()\n",
|
||||
"# base_emb_2 = m(read_img(f\"./data/bottle/test/{t}/001.png\")).detach().numpy()\n",
|
||||
"# base_emb_3 = m(read_img(f\"./data/bottle/test/{t}/002.png\")).detach().numpy()\n",
|
||||
"# base_emb_4 = m(read_img(f\"./data/bottle/test/{t}/003.png\")).detach().numpy()\n",
|
||||
"# base_emb_5 = m(read_img(f\"./data/bottle/test/{t}/004.png\")).detach().numpy()\n",
|
||||
"# base_emb = merge_embeddings([base_emb_1, base_emb_2, base_emb_3, base_emb_4, base_emb_5])\n",
|
||||
"# return base_emb\n",
|
||||
"\n",
|
||||
"MAIN_TYPE=\"bottle\"\n",
|
||||
"\n",
|
||||
"def calc_base_emb(t, nr):\n",
|
||||
" embs = []\n",
|
||||
" for i in range(nr):\n",
|
||||
" if t == \"good\":\n",
|
||||
" emb = m(read_img(f\"./data/{MAIN_TYPE}/train/{t}/{i:03d}.png\")).detach().numpy()\n",
|
||||
" else:\n",
|
||||
" emb = m(read_img(f\"./data/{MAIN_TYPE}/test/{t}/{i:03d}.png\")).detach().numpy()\n",
|
||||
" embs.append(emb)\n",
|
||||
" base_emb = merge_embeddings(embs)\n",
|
||||
" return base_emb\n",
|
||||
"\n",
|
||||
"shotnr=5\n",
|
||||
"goodnr=5\n",
|
||||
"\n",
|
||||
"types = {#\"good\": calc_base_emb(\"good\", goodnr), \n",
|
||||
" #\"bad\": merge_embeddings([calc_base_emb(\"broken_large\", shotnr), calc_base_emb(\"broken_small\", shotnr), calc_base_emb(\"contamination\", shotnr)]),\n",
|
||||
" \"broken_large\": calc_base_emb(\"broken_large\", shotnr), \n",
|
||||
" \"broken_small\": calc_base_emb(\"broken_small\", shotnr), \n",
|
||||
" \"contamination\": calc_base_emb(\"contamination\", shotnr)\n",
|
||||
" }\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"#types = {#\"good\": calc_base_emb(\"good\", goodnr), \n",
|
||||
"# #\"bad\": merge_embeddings([calc_base_emb(\"bent_wire\", shotnr), calc_base_emb(\"cable_swap\", shotnr), calc_base_emb(\"combined\", shotnr), \n",
|
||||
"# # calc_base_emb(\"cut_inner_insulation\", shotnr), calc_base_emb(\"cut_outer_insulation\", shotnr), calc_base_emb(\"missing_cable\", shotnr), \n",
|
||||
"# # calc_base_emb(\"missing_wire\", shotnr), calc_base_emb(\"poke_insulation\", shotnr)]),\n",
|
||||
"# \"bent_wire\": calc_base_emb(\"bent_wire\", shotnr),\n",
|
||||
"# \"cable_swap\": calc_base_emb(\"cable_swap\", shotnr), \n",
|
||||
"# \"combined\": calc_base_emb(\"combined\", shotnr), \n",
|
||||
"# \"cut_inner_insulation\": calc_base_emb(\"cut_inner_insulation\", shotnr), \n",
|
||||
"# \"cut_outer_insulation\": calc_base_emb(\"cut_outer_insulation\", shotnr), \n",
|
||||
"# \"missing_cable\": calc_base_emb(\"missing_cable\", shotnr), \n",
|
||||
"# \"missing_wire\": calc_base_emb(\"missing_wire\", shotnr), \n",
|
||||
"# \"poke_insulation\": calc_base_emb(\"poke_insulation\", shotnr), \n",
|
||||
"# }\n",
|
||||
"\n",
|
||||
"# euclidean distance\n",
|
||||
"euclidean_distance_metric = lambda emb1,emb2 : math.pow(euclidean_distances([emb1.flatten(), emb2.flatten()], [emb1.flatten(), emb2.flatten()])[0][1], 2)\n",
|
||||
"# cosine metric\n",
|
||||
"cosine_similarity_metric = lambda emb1,emb2 : cosine_similarity([emb1.flatten(), emb2.flatten()])[0][1]\n",
|
||||
"\n",
|
||||
"lmnn = LMNN(n_neighbors=2, learn_rate=1e-3, verbose=True)\n",
|
||||
"#lmnn.fit(data, [0,0,0,1,1,1,2,2,2,3,3,3])\n",
|
||||
"\n",
|
||||
"lmnn_similarity_metric = lambda emb1,emb2 : lmnn.get_metric()(emb1.flatten(), emb2.flatten())\n",
|
||||
"\n",
|
||||
"Smaller_Better_Metric = False\n",
|
||||
"\n",
|
||||
"def test(type):\n",
|
||||
" predictions = []\n",
|
||||
"\n",
|
||||
" _, _, files = next(os.walk(f\"./data/{MAIN_TYPE}/test/{type}/\"))\n",
|
||||
" file_count = len(files)\n",
|
||||
" for i in range(5,file_count):\n",
|
||||
" print(\".\", end=\"\")\n",
|
||||
"\n",
|
||||
" emb = m(read_img(f\"./data/{MAIN_TYPE}/test/{type}/{i:03d}.png\")).detach().numpy()\n",
|
||||
" curr_score = .0 if Smaller_Better_Metric else 999999.0\n",
|
||||
" max_type = \"\"\n",
|
||||
" for t in types.keys():\n",
|
||||
"# for t in [\"good\", \"bad\"]:\n",
|
||||
" score = compare_embeddings(emb, types[t], euclidean_distance_metric)\n",
|
||||
" \n",
|
||||
" if Smaller_Better_Metric:\n",
|
||||
" if score > curr_score:\n",
|
||||
" curr_score = score\n",
|
||||
" max_type = t\n",
|
||||
" else:\n",
|
||||
" if score < curr_score:\n",
|
||||
" curr_score = score\n",
|
||||
" max_type = t\n",
|
||||
" pass\n",
|
||||
" predictions.append(max_type)\n",
|
||||
" pass\n",
|
||||
" print()\n",
|
||||
" return np.mean([x == type for x in predictions])\n",
|
||||
"# return np.mean([x == (\"good\" if type == \"good\" else \"bad\") for x in predictions])\n",
|
||||
" pass\n",
|
||||
"\n",
|
||||
"accs = []\n",
|
||||
"for t in types.keys():\n",
|
||||
" if t == \"bad\":\n",
|
||||
" continue\n",
|
||||
" accuracy = test(t)\n",
|
||||
" print(f\"accuracy for {t} = {accuracy}\")\n",
|
||||
" accs.append(accuracy)\n",
|
||||
"print(f\"overall accuracy: {np.mean(accs)}\")\n",
|
||||
"#print(m)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"RESNET 50:\n",
|
||||
"Resulsts:\n",
|
||||
"\n",
|
||||
"bottle:\n",
|
||||
"jeweils 1,3,5 shots normal\n",
|
||||
"[0.5892857142857143 0.7321428571428571 0.75]\n",
|
||||
"\n",
|
||||
"inbalanced - mehr good shots 5,10,15,30 -> alle anderen nur 5\n",
|
||||
"[0.75 0.732 0.696 0.696]\n",
|
||||
"\n",
|
||||
"2 ways nur detektieren ob fehlerhaft oder nicht 1,3,5 shots\n",
|
||||
"[0.8395 0.8315 0.8031]\n",
|
||||
"\n",
|
||||
"inbalance 2 way 5,10,15,30 -> rest 5\n",
|
||||
"[0.8031 0.81893 0.8336 0.8031]\n",
|
||||
"\n",
|
||||
"nur fehlerklasse erkennen 1,3,5\n",
|
||||
"[0.7638 0.7428 0.787]\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"cable:\n",
|
||||
"jeweils 1,3,5 shots normal\n",
|
||||
"[0.21808 0.43815 0.4321478]\n",
|
||||
"\n",
|
||||
"inbalanced - mehr good shots 5,10,15,30 -> alle anderen nur 5\n",
|
||||
"[0.4321478 0.432986 0.42340 0.4464635]\n",
|
||||
"\n",
|
||||
"2 ways nur detektieren ob fehlerhaft oder nicht 1,3,5 shots\n",
|
||||
"[0.8592 0.8772 0.8495]\n",
|
||||
"\n",
|
||||
"inbalance 2 way 5,10,15,30 -> rest 5\n",
|
||||
"[0.8495 0.8180 0.7460 0.6846]\n",
|
||||
"\n",
|
||||
"nur fehlerklasse erkennen 1,3,5\n",
|
||||
"[0.240 0.4740 0.4805]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.14"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
+29
-1
@@ -1 +1,29 @@
|
||||
Test intro
|
||||
\section{Introduction}\label{sec:introduction}
|
||||
\subsection{Motivation}\label{subsec:motivation}
|
||||
For most supervised learning tasks lots of training samples are essential.
|
||||
With too less training data the model will not generalize well and not fit a real world task.
|
||||
Labeling datasets is commonly seen as an expensive task and wants to be avoided as much as possible.\cite{generalAI}
|
||||
That's why there is a machine-learning field called active learning.
|
||||
The general approach is to train a model that predicts within every iteration a ranking metric or Pseudo-Labels which then can be used to rank the importance of samples to be labeled by an oracle.
|
||||
These labeled samples are then used to train the model.\cite{activelearning}
|
||||
|
||||
The goal of this practical work is to test active learning within a simple classification task and evaluate its performance.
|
||||
\subsection{Research Questions}\label{subsec:research-questions}
|
||||
|
||||
\subsubsection{Is Few-Shot learning a suitable fit for anomaly detection?}
|
||||
|
||||
Should Few-Shot learning be used for anomaly detection tasks?
|
||||
How does it compare to well established algorithms such as Patchcore or EfficientAD?
|
||||
|
||||
\subsubsection{How does disbalancing the Shot number affect performance?}
|
||||
Does giving the Few-Shot learner more good than bad samples improve the model performance?
|
||||
|
||||
\subsubsection{How does the 3 methods perform in only detecting the anomaly class?}
|
||||
How much does the performance improve if only detecting an anomaly or not?
|
||||
How does it compare to PatchCore and EfficientAD?
|
||||
|
||||
\subsubsection{Extra: How does Euclidean distance compare to Cosine-similarity when using ResNet as a feature-extractor?}
|
||||
I've tried different distance measures -> but results are pretty much the same.
|
||||
|
||||
\subsection{Outline}\label{subsec:outline}
|
||||
|
||||
Reference in new issue
Block a user