train_colab.ipynb
8,767 bytes
| 1 | { |
|---|---|
| 2 | "cells": [ |
| 3 | { |
| 4 | "cell_type": "markdown", |
| 5 | "id": "intro", |
| 6 | "metadata": {}, |
| 7 | "source": [ |
| 8 | "# CountrySense: treening Google Colabis\n", |
| 9 | "\n", |
| 10 | "Terve toru: andmestiku allalaadimine, puhastamine, treening kahes faasis ja hindamine.\n", |
| 11 | "\n", |
| 12 | "Enne alustamist: **Runtime > Change runtime type > T4 GPU**.\n", |
| 13 | "\n", |
| 14 | "Kogu jooks kestab T4 peal suurusjärgus 2 tundi, millest puhastus on umbes 20 minutit\n", |
| 15 | "ja treening tund kuni poolteist.\n", |
| 16 | "\n", |
| 17 | "**Token.** Repo on privaatne. Tee enne alustamist valmis:\n", |
| 18 | "github.com/settings/personal-access-tokens > Generate new token > ainult `countrysense`\n", |
| 19 | "repo, õigus Contents: Read-only. Tokenit küsitakse jooksu ajal ja seda ei salvestata\n", |
| 20 | "notebooki.\n", |
| 21 | "\n", |
| 22 | "**Salvestamine.** Colabi masin kaob sessiooni lõpus koos kogu `runs/` kaustaga.\n", |
| 23 | "Viimane lahter kopeerib tulemused Google Drive'i. Ära jäta seda vahele.\n", |
| 24 | "\n", |
| 25 | "**Milline variant.** See notebook jookseb `configs/all_countries.yaml` peal, kus\n", |
| 26 | "`min_per_class: 1`, ehk iga riik, mis puhastuse üle elab, jääb klassiks alles.\n", |
| 27 | "Kui tahad selle asemel varasemat 56 klassi varianti, vaheta treeningulahtris konfiks\n", |
| 28 | "`configs/default.yaml` ja paranda allpool kaustanimed `clip_vit_b16_all` -> `clip_vit_b16`." |
| 29 | ] |
| 30 | }, |
| 31 | { |
| 32 | "cell_type": "code", |
| 33 | "execution_count": null, |
| 34 | "id": "gpu", |
| 35 | "metadata": {}, |
| 36 | "outputs": [], |
| 37 | "source": [ |
| 38 | "!nvidia-smi" |
| 39 | ] |
| 40 | }, |
| 41 | { |
| 42 | "cell_type": "code", |
| 43 | "execution_count": null, |
| 44 | "id": "clone", |
| 45 | "metadata": {}, |
| 46 | "outputs": [], |
| 47 | "source": [ |
| 48 | "import getpass, subprocess\n", |
| 49 | "\n", |
| 50 | "# Privaatne repo: token küsitakse jooksu ajal ja jääb ainult mällu. Kloonitav URL\n", |
| 51 | "# ehitatakse eraldi, et token ei satuks Colabi väljundisse ega notebooki faili.\n", |
| 52 | "token = getpass.getpass(\"GitHubi token (Contents: Read-only): \")\n", |
| 53 | "subprocess.run(\n", |
| 54 | " [\"git\", \"clone\", f\"https://{token}@github.com/rasmusjy/countrysense.git\"],\n", |
| 55 | " check=True, capture_output=True,\n", |
| 56 | ")\n", |
| 57 | "del token\n", |
| 58 | "print(\"Kloonitud.\")" |
| 59 | ] |
| 60 | }, |
| 61 | { |
| 62 | "cell_type": "code", |
| 63 | "execution_count": null, |
| 64 | "id": "setup", |
| 65 | "metadata": {}, |
| 66 | "outputs": [], |
| 67 | "source": [ |
| 68 | "%cd countrysense\n", |
| 69 | "# Kloonitud URL sisaldab tokenit, seega eemaldame selle .git/config'ist kohe ära.\n", |
| 70 | "!git remote set-url origin https://github.com/rasmusjy/countrysense.git\n", |
| 71 | "!pip install -q -r requirements.txt kagglehub" |
| 72 | ] |
| 73 | }, |
| 74 | { |
| 75 | "cell_type": "markdown", |
| 76 | "id": "md-data", |
| 77 | "metadata": {}, |
| 78 | "source": [ |
| 79 | "## Andmestik\n", |
| 80 | "\n", |
| 81 | "Kaggle andmestik `ubitquitin/geolocation-geoguessr-images-50k`: umbes 50 000 Street View\n", |
| 82 | "pilti, kaustad riikide kaupa, 6,7 GB. Allalaadimine kagglehub kaudu ei nõua Kaggle API võtit." |
| 83 | ] |
| 84 | }, |
| 85 | { |
| 86 | "cell_type": "code", |
| 87 | "execution_count": null, |
| 88 | "id": "download", |
| 89 | "metadata": {}, |
| 90 | "outputs": [], |
| 91 | "source": [ |
| 92 | "import kagglehub\n", |
| 93 | "from pathlib import Path\n", |
| 94 | "\n", |
| 95 | "ds_path = Path(kagglehub.dataset_download(\"ubitquitin/geolocation-geoguessr-images-50k\"))\n", |
| 96 | "\n", |
| 97 | "# Leiame kausta, mille alamkaustad on riigid (paigutus voib versiooniti erineda).\n", |
| 98 | "candidates = [p for p in [ds_path, *ds_path.rglob(\"*\")] if p.is_dir()]\n", |
| 99 | "raw_dir = max(candidates, key=lambda p: sum(1 for c in p.iterdir() if c.is_dir()))\n", |
| 100 | "print(\"Toorandmed:\", raw_dir)\n", |
| 101 | "print(\"Riikide kaustu:\", sum(1 for c in raw_dir.iterdir() if c.is_dir()))" |
| 102 | ] |
| 103 | }, |
| 104 | { |
| 105 | "cell_type": "markdown", |
| 106 | "id": "md-clean", |
| 107 | "metadata": {}, |
| 108 | "source": [ |
| 109 | "## Puhastamine\n", |
| 110 | "\n", |
| 111 | "Eemaldab rikutud, liiga väikesed, tumedad, heledad ja udused pildid ning phash-duplikaadid.\n", |
| 112 | "`--min-per-class 1` hoiab kõik riigid alles; kvaliteedikontroll käib ikka.\n", |
| 113 | "Kirjutab `data/manifest.csv` ja raporti `data/cleaning_report.csv`." |
| 114 | ] |
| 115 | }, |
| 116 | { |
| 117 | "cell_type": "code", |
| 118 | "execution_count": null, |
| 119 | "id": "clean", |
| 120 | "metadata": {}, |
| 121 | "outputs": [], |
| 122 | "source": [ |
| 123 | "!python -m countrysense.clean --raw-dir \"{raw_dir}\" --min-per-class 1" |
| 124 | ] |
| 125 | }, |
| 126 | { |
| 127 | "cell_type": "code", |
| 128 | "execution_count": null, |
| 129 | "id": "dist", |
| 130 | "metadata": {}, |
| 131 | "outputs": [], |
| 132 | "source": [ |
| 133 | "# Kui palju pilte riigi kohta puhastuse järel alles on. Määrab, kui palju klasse\n", |
| 134 | "# on sisuliselt õpitavad ja kui palju jääb käputäie näite peale.\n", |
| 135 | "import pandas as pd\n", |
| 136 | "\n", |
| 137 | "counts = pd.read_csv(\"data/manifest.csv\")[\"country\"].value_counts()\n", |
| 138 | "print(f\"Riike kokku: {len(counts)}, pilte: {int(counts.sum())}\")\n", |
| 139 | "for lo in (100, 50, 20, 10, 5, 1):\n", |
| 140 | " print(f\" vähemalt {lo:>3} pilti: {int((counts >= lo).sum()):>3} riiki\")\n", |
| 141 | "print(\"\\nKõige väiksemad:\")\n", |
| 142 | "print(counts.tail(12).to_string())" |
| 143 | ] |
| 144 | }, |
| 145 | { |
| 146 | "cell_type": "markdown", |
| 147 | "id": "md-train", |
| 148 | "metadata": {}, |
| 149 | "source": [ |
| 150 | "## Treening\n", |
| 151 | "\n", |
| 152 | "Faas 1: põhivõrk külmutatud, treenime ainult klassifitseerimispead (3 epohhi).\n", |
| 153 | "Faas 2: viimased 4 transformeri plokki lahti, madal õppemäär (8 epohhi).\n", |
| 154 | "\n", |
| 155 | "Parim checkpoint salvestatakse faili `runs/clip_vit_b16_all/best.pt`." |
| 156 | ] |
| 157 | }, |
| 158 | { |
| 159 | "cell_type": "code", |
| 160 | "execution_count": null, |
| 161 | "id": "train", |
| 162 | "metadata": {}, |
| 163 | "outputs": [], |
| 164 | "source": [ |
| 165 | "!python -m countrysense.train --config configs/all_countries.yaml" |
| 166 | ] |
| 167 | }, |
| 168 | { |
| 169 | "cell_type": "markdown", |
| 170 | "id": "md-eval", |
| 171 | "metadata": {}, |
| 172 | "source": [ |
| 173 | "## Hindamine testihulgal\n", |
| 174 | "\n", |
| 175 | "Makro-F1 keskmistab üle klasside, mitte piltide, seega karistab see väikeseid\n", |
| 176 | "riike täie raskusega. Top-1 langeb vähem, sest testihulka valitsevad suured riigid.\n", |
| 177 | "Vahe nende kahe vahel ütlebki, kui palju katvus täpsusest võttis." |
| 178 | ] |
| 179 | }, |
| 180 | { |
| 181 | "cell_type": "code", |
| 182 | "execution_count": null, |
| 183 | "id": "eval", |
| 184 | "metadata": {}, |
| 185 | "outputs": [], |
| 186 | "source": [ |
| 187 | "!python -m countrysense.evaluate --checkpoint runs/clip_vit_b16_all/best.pt" |
| 188 | ] |
| 189 | }, |
| 190 | { |
| 191 | "cell_type": "code", |
| 192 | "execution_count": null, |
| 193 | "id": "cm", |
| 194 | "metadata": {}, |
| 195 | "outputs": [], |
| 196 | "source": [ |
| 197 | "from IPython.display import Image as IPyImage\n", |
| 198 | "\n", |
| 199 | "IPyImage(\"runs/clip_vit_b16_all/confusion_matrix.png\")" |
| 200 | ] |
| 201 | }, |
| 202 | { |
| 203 | "cell_type": "markdown", |
| 204 | "id": "md-try", |
| 205 | "metadata": {}, |
| 206 | "source": [ |
| 207 | "## Proovi ühe pildiga" |
| 208 | ] |
| 209 | }, |
| 210 | { |
| 211 | "cell_type": "code", |
| 212 | "execution_count": null, |
| 213 | "id": "try", |
| 214 | "metadata": {}, |
| 215 | "outputs": [], |
| 216 | "source": [ |
| 217 | "import subprocess\n", |
| 218 | "import sys\n", |
| 219 | "\n", |
| 220 | "import pandas as pd\n", |
| 221 | "\n", |
| 222 | "# subprocess, mitte ! -maagia: riigikaustade nimedes on tühikuid ja \"{...}\"\n", |
| 223 | "# asendus shellikäsu sees komistab nende otsa.\n", |
| 224 | "sample = pd.read_csv(\"runs/clip_vit_b16_all/test.csv\").sample(1).iloc[0]\n", |
| 225 | "print(\"Tegelik riik:\", sample[\"country\"], \"\\n\")\n", |
| 226 | "\n", |
| 227 | "done = subprocess.run(\n", |
| 228 | " [sys.executable, \"-m\", \"countrysense.predict\",\n", |
| 229 | " \"--checkpoint\", \"runs/clip_vit_b16_all/best.pt\",\n", |
| 230 | " \"--image\", sample[\"path\"]],\n", |
| 231 | " capture_output=True, text=True,\n", |
| 232 | ")\n", |
| 233 | "print(done.stdout or done.stderr)" |
| 234 | ] |
| 235 | }, |
| 236 | { |
| 237 | "cell_type": "markdown", |
| 238 | "id": "md-save", |
| 239 | "metadata": {}, |
| 240 | "source": [ |
| 241 | "## Salvestamine Google Drive'i\n", |
| 242 | "\n", |
| 243 | "Ära jäta seda vahele. Colabi sessiooni lõppedes kustub `runs/` koos kõigega, mis\n", |
| 244 | "treening tootis. Läheb kausta `MyDrive/countrysense-all`, eraldi varasemast\n", |
| 245 | "56 klassi jooksust, et mõlemad jääksid alles ja neid saaks kõrvuti võrrelda." |
| 246 | ] |
| 247 | }, |
| 248 | { |
| 249 | "cell_type": "code", |
| 250 | "execution_count": null, |
| 251 | "id": "save", |
| 252 | "metadata": {}, |
| 253 | "outputs": [], |
| 254 | "source": [ |
| 255 | "from google.colab import drive\n", |
| 256 | "\n", |
| 257 | "drive.mount(\"/content/drive\")\n", |
| 258 | "\n", |
| 259 | "DEST = \"/content/drive/MyDrive/countrysense-all\"\n", |
| 260 | "\n", |
| 261 | "# best.pt on demo jaoks, ülejäänud selleks, et tulemusi saaks hiljem tsiteerida.\n", |
| 262 | "!mkdir -p \"{DEST}\"\n", |
| 263 | "!cp runs/clip_vit_b16_all/best.pt \"{DEST}/best.pt\"\n", |
| 264 | "!cp -f runs/clip_vit_b16_all/*.png \"{DEST}/\"\n", |
| 265 | "!cp -f runs/clip_vit_b16_all/*.csv \"{DEST}/\"\n", |
| 266 | "!cp -f runs/clip_vit_b16_all/*.json \"{DEST}/\"\n", |
| 267 | "!cp -f data/cleaning_report.csv \"{DEST}/\"\n", |
| 268 | "!ls -lh \"{DEST}/\"" |
| 269 | ] |
| 270 | } |
| 271 | ], |
| 272 | "metadata": { |
| 273 | "accelerator": "GPU", |
| 274 | "colab": { |
| 275 | "gpuType": "T4", |
| 276 | "provenance": [] |
| 277 | }, |
| 278 | "kernelspec": { |
| 279 | "display_name": "Python 3", |
| 280 | "name": "python3" |
| 281 | }, |
| 282 | "language_info": { |
| 283 | "name": "python" |
| 284 | } |
| 285 | }, |
| 286 | "nbformat": 4, |
| 287 | "nbformat_minor": 5 |
| 288 | } |
| 289 | |