profileShare

rasmusjy / countrysense

Read-only snapshot

No repository description.

main default branch 19 files Expires Sep 13, 2026, 9:06 AM

Commit

Add the all-countries coverage variant and keep small classes trainable

Dropping the minimum-images-per-class threshold brings back the countries that
the default config filters out. That only helps if those classes actually get
trained on, so the stratified split now always leaves at least one image in
train and lets val and test take what is left over. A class with no training
example would appear in the output without having been learned, which is worse
than a missing measurement. Large classes split as before.

The trade-off is written down in configs/all_countries.yaml: the default
threshold of 100 keeps 56 countries at 73.9% top-1, and lowering it adds
countries whose macro F1 will pull the average down.
commit 5856546

4 changed files with +195 and −46

Jump to a changed file
  1. .gitignore +3 −0
  2. configs/all_countries.yaml +35 −0
  3. countrysense/data.py +9 −2
  4. notebooks/train_colab.ipynb +148 −44
modified .gitignore +3 −0
@@ -8,3 +8,6 @@runs/
8 8 .venv/
9 9 venv/
10 10 .DS_Store
11 +
12 +# Gradio runtime cache, not part of the project.
13 +.gradio/
added configs/all_countries.yaml +35 −0
@@ -0,0 +1,35 @@
1 +# Katvuse variant: iga riik, mis puhastuse üle elab, jääb klassiks alles.
2 +#
3 +# Vahetuskaup on teadlik. Lävendiga 100 jäi alles 56 riiki ja test andis top-1
4 +# 73,9%. Lävendiga 1 tuleb juurde kümneid riike, millel on käputäis pilti; need
5 +# klassid jäävad tõenäoliselt nulli ligi ja makro-F1 langeb, sest see keskmistab
6 +# üle klasside, mitte piltide. Top-1 langeb vähem, sest testihulka valitsevad
7 +# endiselt suured riigid, aga langeb ikka: iga uus klass on suurtele riikidele
8 +# üks lisavõimalus eksida.
9 +#
10 +# Väljund läheb eraldi kausta, et 56 klassi tulemus jääks kõrvale alles ja neid
11 +# saaks ausalt kõrvuti näidata.
12 +data:
13 + raw_dir: data/raw # iga alamkaust on üks riik
14 + manifest: data/manifest.csv # clean.py väljund; kui puudub, skannitakse raw_dir
15 + min_per_class: 1 # kõik riigid alles; split hoolitseb, et igaühel oleks treeningpilt
16 + val_frac: 0.1
17 + test_frac: 0.1
18 + num_workers: 2
19 +
20 +model:
21 + backbone: vit_base_patch16_clip_224.openai
22 + dropout: 0.2
23 +
24 +train:
25 + seed: 42
26 + batch_size: 64
27 + epochs_head: 3 # faas 1: ainult klassifitseerimispea
28 + epochs_finetune: 8 # faas 2: viimased plokid lahti
29 + unfreeze_blocks: 4
30 + lr_head: 1.0e-3
31 + lr_head_finetune: 1.0e-4
32 + lr_backbone: 2.0e-5
33 + weight_decay: 0.01
34 + label_smoothing: 0.1
35 + out_dir: runs/clip_vit_b16_all
modified countrysense/data.py +9 −2
@@ -23,11 +23,18 @@def scan_image_folder(root):
23 23
24 24 def stratified_split(df, val_frac=0.1, test_frac=0.1, seed=42):
25 25 # Jagame iga riigi sees eraldi, et val/test sisaldaks kõiki klasse.
26 + #
27 + # Väikese riigi puhul on treeningpilt tähtsam kui hindamispilt: klass, mille
28 + # kohta ei ole ühtegi treeningnäidet, on mudeli väljundis olemas, aga õpitud
29 + # ei ole, ja see on hullem kui puuduv mõõtmine. Seepärast jääb alati
30 + # vähemalt üks pilt treeningusse ning val ja test saavad ainult selle, mis
31 + # üle jääb. Suurte klasside jaotus ei muutu.
26 32 parts = {"train": [], "val": [], "test": []}
27 33 for _, group in df.groupby("country"):
28 34 group = group.sample(frac=1.0, random_state=seed)
29 - n_val = max(1, round(len(group) * val_frac))
30 - n_test = max(1, round(len(group) * test_frac))
35 + spare = len(group) - 1
36 + n_val = min(max(1, round(len(group) * val_frac)), spare)
37 + n_test = min(max(1, round(len(group) * test_frac)), spare - n_val)
31 38 parts["val"].append(group.iloc[:n_val])
32 39 parts["test"].append(group.iloc[n_val:n_val + n_test])
33 40 parts["train"].append(group.iloc[n_val + n_test:])
modified notebooks/train_colab.ipynb +148 −44
@@ -1,68 +1,92 @@
1 1 {
2 - "nbformat": 4,
3 - "nbformat_minor": 5,
4 - "metadata": {
5 - "accelerator": "GPU",
6 - "colab": {
7 - "provenance": [],
8 - "gpuType": "T4"
9 - },
10 - "kernelspec": {
11 - "name": "python3",
12 - "display_name": "Python 3"
13 - },
14 - "language_info": {
15 - "name": "python"
16 - }
17 - },
18 2 "cells": [
19 3 {
20 4 "cell_type": "markdown",
5 + "id": "intro",
21 6 "metadata": {},
22 7 "source": [
23 8 "# CountrySense: treening Google Colabis\n",
24 9 "\n",
25 10 "Terve toru: andmestiku allalaadimine, puhastamine, treening kahes faasis ja hindamine.\n",
26 11 "\n",
27 - "Enne alustamist: Runtime > Change runtime type > T4 GPU.\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",
28 21 "\n",
29 - "Kogu treening kestab T4 peal suurusjärgus 1 kuni 2 tundi."
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`."
30 29 ]
31 30 },
32 31 {
33 32 "cell_type": "code",
34 - "metadata": {},
35 33 "execution_count": null,
34 + "id": "gpu",
35 + "metadata": {},
36 36 "outputs": [],
37 37 "source": [
38 38 "!nvidia-smi"
39 39 ]
40 40 },
41 41 {
42 42 "cell_type": "code",
43 + "execution_count": null,
44 + "id": "clone",
43 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",
44 63 "execution_count": null,
64 + "id": "setup",
65 + "metadata": {},
45 66 "outputs": [],
46 67 "source": [
47 - "# Asenda URL oma repoga.\n",
48 - "!git clone https://github.com/rasmju/countrysense.git\n",
49 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",
50 71 "!pip install -q -r requirements.txt kagglehub"
51 72 ]
52 73 },
53 74 {
54 75 "cell_type": "markdown",
76 + "id": "md-data",
55 77 "metadata": {},
56 78 "source": [
57 79 "## Andmestik\n",
58 80 "\n",
59 - "Kaggle andmestik `ubitquitin/geolocation-geoguessr-images-50k`: umbes 50 000 Street View pilti, kaustad riikide kaupa. Allalaadimine kagglehub kaudu ei nõua Kaggle API võtit."
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."
60 83 ]
61 84 },
62 85 {
63 86 "cell_type": "code",
64 - "metadata": {},
65 87 "execution_count": null,
88 + "id": "download",
89 + "metadata": {},
66 90 "outputs": [],
67 91 "source": [
68 92 "import kagglehub\n",
@@ -79,106 +103,186 @@
79 103 },
80 104 {
81 105 "cell_type": "markdown",
106 + "id": "md-clean",
82 107 "metadata": {},
83 108 "source": [
84 109 "## Puhastamine\n",
85 110 "\n",
86 - "Eemaldab rikutud, liiga väikesed, tumedad, heledad ja udused pildid ning phash-duplikaadid. Riigid, millel on alla 100 pildi, jäetakse välja. Kirjutab `data/manifest.csv` ja raporti `data/cleaning_report.csv`."
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`."
87 114 ]
88 115 },
89 116 {
90 117 "cell_type": "code",
118 + "execution_count": null,
119 + "id": "clean",
91 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",
92 128 "execution_count": null,
129 + "id": "dist",
130 + "metadata": {},
93 131 "outputs": [],
94 132 "source": [
95 - "!python -m countrysense.clean --raw-dir \"{raw_dir}\""
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())"
96 143 ]
97 144 },
98 145 {
99 146 "cell_type": "markdown",
147 + "id": "md-train",
100 148 "metadata": {},
101 149 "source": [
102 150 "## Treening\n",
103 151 "\n",
104 - "Faas 1: backbone külmutatud, treenime ainult klassifitseerimispead (3 epohhi).\n",
152 + "Faas 1: põhivõrk külmutatud, treenime ainult klassifitseerimispead (3 epohhi).\n",
105 153 "Faas 2: viimased 4 transformeri plokki lahti, madal õppemäär (8 epohhi).\n",
106 154 "\n",
107 - "Parim checkpoint salvestatakse faili `runs/clip_vit_b16/best.pt`."
155 + "Parim checkpoint salvestatakse faili `runs/clip_vit_b16_all/best.pt`."
108 156 ]
109 157 },
110 158 {
111 159 "cell_type": "code",
112 - "metadata": {},
113 160 "execution_count": null,
161 + "id": "train",
162 + "metadata": {},
114 163 "outputs": [],
115 164 "source": [
116 - "!python -m countrysense.train --config configs/default.yaml"
165 + "!python -m countrysense.train --config configs/all_countries.yaml"
117 166 ]
118 167 },
119 168 {
120 169 "cell_type": "markdown",
170 + "id": "md-eval",
121 171 "metadata": {},
122 172 "source": [
123 - "## Hindamine testihulgal"
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."
124 178 ]
125 179 },
126 180 {
127 181 "cell_type": "code",
128 - "metadata": {},
129 182 "execution_count": null,
183 + "id": "eval",
184 + "metadata": {},
130 185 "outputs": [],
131 186 "source": [
132 - "!python -m countrysense.evaluate --checkpoint runs/clip_vit_b16/best.pt"
187 + "!python -m countrysense.evaluate --checkpoint runs/clip_vit_b16_all/best.pt"
133 188 ]
134 189 },
135 190 {
136 191 "cell_type": "code",
137 - "metadata": {},
138 192 "execution_count": null,
193 + "id": "cm",
194 + "metadata": {},
139 195 "outputs": [],
140 196 "source": [
141 197 "from IPython.display import Image as IPyImage\n",
142 - "IPyImage(\"runs/clip_vit_b16/confusion_matrix.png\")"
198 + "\n",
199 + "IPyImage(\"runs/clip_vit_b16_all/confusion_matrix.png\")"
143 200 ]
144 201 },
145 202 {
146 203 "cell_type": "markdown",
204 + "id": "md-try",
147 205 "metadata": {},
148 206 "source": [
149 207 "## Proovi ühe pildiga"
150 208 ]
151 209 },
152 210 {
153 211 "cell_type": "code",
154 - "metadata": {},
155 212 "execution_count": null,
213 + "id": "try",
214 + "metadata": {},
156 215 "outputs": [],
157 216 "source": [
217 + "import subprocess\n",
218 + "import sys\n",
219 + "\n",
158 220 "import pandas as pd\n",
159 221 "\n",
160 - "sample = pd.read_csv(\"runs/clip_vit_b16/test.csv\").sample(1).iloc[0]\n",
161 - "print(\"Tegelik riik:\", sample[\"country\"])\n",
162 - "!python -m countrysense.predict --image \"{sample['path']}\""
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)"
163 234 ]
164 235 },
165 236 {
166 237 "cell_type": "markdown",
238 + "id": "md-save",
167 239 "metadata": {},
168 240 "source": [
169 - "## Checkpointi salvestamine Google Drive'i (valikuline)"
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."
170 246 ]
171 247 },
172 248 {
173 249 "cell_type": "code",
174 - "metadata": {},
175 250 "execution_count": null,
251 + "id": "save",
252 + "metadata": {},
176 253 "outputs": [],
177 254 "source": [
178 - "# from google.colab import drive\n",
179 - "# drive.mount(\"/content/drive\")\n",
180 - "# !cp runs/clip_vit_b16/best.pt /content/drive/MyDrive/countrysense_best.pt"
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}/\""
181 269 ]
182 270 }
183 - ]
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
184 288 }