profileShare

rasmusjy / countrysense

Read-only snapshot

No repository description.

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

Commit

Add CountrySense: CLIP ViT country classifier, training pipeline, docs

commit fda842f

15 changed files with +1377 and −0

Jump to a changed file
  1. .gitignore +10 −0
  2. README.et.md +113 −0
  3. README.md +65 −0
  4. app.py +35 −0
  5. configs/default.yaml +24 −0
  6. countrysense/__init__.py +0 −0
  7. countrysense/clean.py +105 −0
  8. countrysense/data.py +66 −0
  9. countrysense/evaluate.py +99 −0
  10. countrysense/model.py +48 −0
  11. countrysense/predict.py +33 −0
  12. countrysense/train.py +177 −0
  13. docs/mudel.html +403 −0
  14. notebooks/train_colab.ipynb +184 −0
  15. requirements.txt +15 −0
added .gitignore +10 −0
@@ -0,0 +1,10 @@
1 +__pycache__/
2 +*.pyc
3 +.ipynb_checkpoints/
4 +data/
5 +runs/
6 +*.pt
7 +*.ckpt
8 +.venv/
9 +venv/
10 +.DS_Store
added README.et.md +113 −0
@@ -0,0 +1,113 @@
1 +# CountrySense
2 +
3 +Mudel, mis arvab ühe tänavapildi järgi ära riigi. Peenhäälestatud CLIP-i nägemistransformer (ViT-B/16) klassifitseerib pildi umbes saja riigi vahel. Ainult riigi tase: ei mingeid koordinaate ega linnade äraarvamist.
4 +
5 +Põhjalik arhitektuuri selgitus (kihthaaval, koos andmete valiku ja puhastamisega) on failis [docs/mudel.html](docs/mudel.html). Ava see brauseris.
6 +
7 +## Miks riigi tase, mitte koordinaadid
8 +
9 +Täpne geolokatsioon (laius- ja pikkuskraadi ennustamine) on teadusartikli tasemel ülesanne: vaja on miljoneid pilte, suuri mudeleid ja arvutusressurssi, mida sel projektil ei ole. Riigi klassifitseerimine umbes saja riigi vahel on seevastu tavaline juhendatud õppe ülesanne: eeltreenitud nägemismudeli peenhäälestus mahub tasuta Colabi või Kaggle'i GPU peale ja annab demos üllatavalt hea täpsuse.
10 +
11 +Aus teekaart on selline: kõigepealt riik, hiljem regioon riigi sees (näiteks "Eesti, Lõuna-Eesti"). Koordinaadid ei ole plaanis.
12 +
13 +## Andmestik ja valikukriteeriumid
14 +
15 +Vaikimisi andmestik on Kaggle'i [GeoLocation, Geoguessr Images 50K](https://www.kaggle.com/datasets/ubitquitin/geolocation-geoguessr-images-50k): umbes 50 000 Street View pilti, kaustad riikide kaupa.
16 +
17 +Valikukriteeriumid olid:
18 +
19 +- Sildid on kaustastruktuuris olemas, käsitsi märgendamist pole vaja.
20 +- Umbes 150 riiki, millest pärast puhastust jääb alles 100 ringis. Piisavalt, et ülesanne oleks huvitav, ja piisavalt vähe, et tasuta GPU-ga hakkama saada.
21 +- Suurus (mõni GB) mahub Colabi kettale ja laadimisaeg on talutav.
22 +- Vaba juurdepääs kagglehub kaudu, API võtit pole vaja.
23 +
24 +Kood ei sõltu konkreetsest andmestikust: sobib iga kaust paigutusega `root/<riik>/<pilt>`. Suuremaks skaleerimiseks on hea kandidaat OpenStreetView-5M alamhulk.
25 +
26 +Oluline ausus: Street View kaamera katvus on kaldu rikaste riikide poole ja mudel õpib paratamatult ka kaamera artefakte (põlvkond, resolutsioon), mitte ainult maastikku. See on osa sellest, miks demo tundub "maagiline", ja see on dokumenteeritud, mitte maha vaikitud.
27 +
28 +## Andmete puhastamine
29 +
30 +`countrysense/clean.py` käib kõik pildid läbi ja kirjutab manifesti (`data/manifest.csv`) ning raporti (`data/cleaning_report.csv`), kus iga välja jäetud pildi juures on põhjus.
31 +
32 +| Samm | Reegel | Miks |
33 +|---|---|---|
34 +| Rikutud failid | PIL ei suuda avada | Katkine allalaadimine kukutaks treeningu |
35 +| Suurus | lühem külg < 128 px | Liiga vähe infot, ViT sisend on 224 px |
36 +| Heledus | keskmine < 20 või > 235 | Praktiliselt mustad või läbipõlenud kaadrid |
37 +| Teravus | gradiendi energia < 25 | Udused kaadrid (liikumine, vihm objektiivil) |
38 +| Duplikaadid | sama phash (perceptual hash) | Street View annab kattuvaid kaadreid; duplikaat train- ja testihulgas oleks leke |
39 +| Väikesed klassid | riigil < 100 pilti | Liiga vähe, et õppida ja usaldusväärselt mõõta |
40 +| Suured klassid | lagi 2000 pilti riigi kohta | Vähendab tasakaalustamatust juba andmete tasemel |
41 +
42 +Läved on käsurea argumendid, vaikeväärtused on tabelis.
43 +
44 +## Mudel ja valikukriteeriumid
45 +
46 +Backbone on OpenAI CLIP-i eeltreenitud ViT-B/16 (`vit_base_patch16_clip_224.openai` timm-i kaudu), 86 miljonit parameetrit. Pea on `Dropout(0.2)` pluss üks `Linear(768 -> riikide arv)` kiht.
47 +
48 +Kaalutud variandid:
49 +
50 +| Variant | Hinnang |
51 +|---|---|
52 +| CLIP ViT-B/16 (valitud) | Eeltreenitud 400M pilt-tekst paari peal, tunnused kodeerivad juba silte, taimestikku, arhitektuuri ja teemärgistust. Lineaarne pea CLIP-i tunnuste peal on tugev geolokatsiooni baastase. Mahub T4 peale. |
53 +| ImageNet ViT/ConvNeXt | Töötab, aga ImageNeti klassid (koeratõud, esemed) on geograafiast kaugemal, transfer on nõrgem. |
54 +| StreetCLIP (ViT-L/14) | Geolokatsiooniks juba häälestatud, aga kolm korda suurem, peenhäälestus tasuta T4 peal on kitsas. Hea järgmine samm, kui riistvara lubab. |
55 +| Nullist treenitud CNN | Selle andmemahuga lootusetu, eeltreening on kogu projekti võimaldaja. |
56 +
57 +Miks ainult üks lineaarne kiht peas: CLIP-i tunnused on juba suures osas lineaarselt eraldatavad ja sügavam pea kipub selle andmemahuga üle sobituma. Detailne selgitus, mida transformeri sees olevad dense-kihid teevad, on failis [docs/mudel.html](docs/mudel.html).
58 +
59 +## Treening
60 +
61 +Kaks faasi, kokku umbes 1 kuni 2 tundi tasuta T4 GPU peal:
62 +
63 +1. **Faas 1, lineaarne sondeerimine (3 epohhi, lr 1e-3).** Backbone on külmutatud, treenime ainult pead. Juhuslikult initsialiseeritud pea gradiendid on alguses suured ja rikuksid eeltreenitud kaalud, seega laseme peal enne stabiliseeruda.
64 +2. **Faas 2, osaline peenhäälestus (8 epohhi, lr 2e-5 backbone / 1e-4 pea, koosinusgraafik).** Avame ainult viimased 4 transformeri plokki ja lõpu normi. Alumised kihid õpivad üldisi servi ja tekstuure, mis on niigi head; ülemised kihid kannavad semantikat, mida tasub geograafiale kohandada. Vähem avatud kihte tähendab ka vähem mälu ja väiksemat katastroofilise unustamise riski.
65 +
66 +Muud valikud: cross-entropy koos label smoothing 0.1 (piirialade sildid on loomupäraselt mürarikkad), AdamW, AMP (poolttäpsus), gradientide kärpimine, `WeightedRandomSampler` klasside tasakaalustamiseks (muidu domineeriksid suurte piltide arvuga riigid nagu USA).
67 +
68 +Augmentatsioonid: `RandomResizedCrop` (kaugus ja kadreering varieeruvad päriselt), kerge `ColorJitter` (valgustus ja kaamera). Teadlikult **ei kasuta** horisontaalset peegeldust: vasak- või parempoolne liiklus on päris geograafiline tunnus, mille peegeldamine õpetaks mudelile valet. Samal põhjusel ei pöörata pilti: horisondi asend on info.
69 +
70 +## Mõõtmine
71 +
72 +`evaluate.py` raporteerib top-1, top-5 ja makro-F1, kirjutab riigipõhise täpsuse tabeli ja confusion matrix'i. Top-5 on geograafias aus mõõdik: Eesti ja Läti segiajamine on palju väiksem viga kui Eesti ja Brasiilia oma, ja vigade seas domineerivadki naabrid (Balti riigid omavahel, Skandinaavia, hispaaniakeelne Ladina-Ameerika).
73 +
74 +## Käivitamine
75 +
76 +Colab (soovitatav): ava `notebooks/train_colab.ipynb`, vali T4 GPU ja käivita kõik lahtrid.
77 +
78 +Kohalikult:
79 +
80 +```bash
81 +pip install -r requirements.txt
82 +python -m countrysense.clean --raw-dir data/raw
83 +python -m countrysense.train --config configs/default.yaml
84 +python -m countrysense.evaluate --checkpoint runs/clip_vit_b16/best.pt
85 +python -m countrysense.predict --image minu_pilt.jpg
86 +```
87 +
88 +Gradio demo: `pip install gradio` ja `python app.py`.
89 +
90 +## Tulemused
91 +
92 +Täida pärast esimest treeningut (`runs/clip_vit_b16/metrics.json`):
93 +
94 +| Mõõdik | Väärtus |
95 +|---|---|
96 +| Top-1 | ... |
97 +| Top-5 | ... |
98 +| Makro-F1 | ... |
99 +
100 +Tüüpiline ootus selle retseptiga ja ~100 riigiga: top-1 vahemikus 50 kuni 70 protsenti ja top-5 üle 80 protsendi, sõltuvalt puhastusest ja epohhide arvust. Juhuslik pakkumine oleks 1 protsent.
101 +
102 +## Piirangud
103 +
104 +- Street View katvus ja kaamera bias: mudel õpib osalt kaamerat, mitte maastikku.
105 +- Riigid, mida andmestikus pole, ennustatakse paratamatult valesti mõneks naabriks.
106 +- Sisetingimustes, dokumentidel ja inimestel pole mudelil mõtet, see on tänavavaadete klassifitseerija.
107 +
108 +## Teekaart
109 +
110 +- [ ] Regioonitase riigi sees (hierarhiline pea: enne riik, siis regioon)
111 +- [ ] Suurem treeningandmestik (OpenStreetView-5M alamhulk)
112 +- [ ] Kalibratsioon ja keeldumine ("pole piisavalt kindel, et pakkuda")
113 +- Plaanis ei ole: täpsed koordinaadid. See on teadustöö territoorium (PIGEON, GeoCLIP) ja nõuab arvutusmahtu, mida see projekt teadlikult väldib.
added README.md +65 −0
@@ -0,0 +1,65 @@
1 +# CountrySense
2 +
3 +Guess the country from a single street-level photo. A fine-tuned CLIP vision transformer classifies images into ~100+ countries. Country classification only: no coordinates, no city guessing.
4 +
5 +Estonian documentation: [README.et.md](README.et.md) (full write-up) and [docs/mudel.html](docs/mudel.html) (layer-by-layer explanation of the architecture, data selection and cleaning).
6 +
7 +## Why country-level, not coordinates
8 +
9 +Exact geolocation (predicting latitude and longitude) is research-grade work: it needs millions of images, large models and compute budgets this project does not have. Country classification over ~100 countries is a tractable supervised problem: fine-tune a pretrained vision model on a free Colab or Kaggle GPU in an afternoon and get accuracy that feels surprising in a demo. Region-level prediction within a country is the honest next step on the roadmap. Coordinates are out of scope.
10 +
11 +## Model
12 +
13 +- Backbone: ViT-B/16 pretrained by OpenAI CLIP (`vit_base_patch16_clip_224.openai` via timm), 86M parameters.
14 +- Head: `Dropout(0.2)` + a single `Linear(768 -> num_countries)` layer.
15 +- Why CLIP: it was pretrained on 400M image-text pairs from the web, so its features already encode things like signage, vegetation, architecture and road furniture. A linear probe on CLIP features is a strong geolocation baseline; fine-tuning the last blocks improves on it.
16 +
17 +Training recipe (fits on a free T4, roughly 1 to 2 hours end to end):
18 +
19 +1. Phase 1, linear probe: backbone frozen, train only the head, 3 epochs, lr 1e-3. This protects pretrained weights from gradients of a randomly initialized head.
20 +2. Phase 2, partial fine-tune: unfreeze the last 4 transformer blocks and the final norm, lr 2e-5 (backbone) / 1e-4 (head), cosine schedule, 8 epochs.
21 +
22 +Cross-entropy with label smoothing 0.1, AdamW, AMP, gradient clipping, `WeightedRandomSampler` for class imbalance. No horizontal flip augmentation: driving side (left vs right traffic) is a real geographic cue that flipping would corrupt.
23 +
24 +## Data
25 +
26 +Default dataset: [GeoLocation, Geoguessr Images 50K](https://www.kaggle.com/datasets/ubitquitin/geolocation-geoguessr-images-50k) on Kaggle (~50k Street View images in per-country folders). Any dataset laid out as `root/<country>/<image>` works.
27 +
28 +Cleaning (`countrysense/clean.py`) drops corrupt files, images smaller than 128px, too dark / too bright / blurry images, perceptual-hash duplicates, and countries with fewer than 100 images. It writes a manifest CSV plus a per-image report of what was dropped and why.
29 +
30 +## Quickstart
31 +
32 +Colab (recommended): open `notebooks/train_colab.ipynb`, select a T4 GPU runtime and run all cells.
33 +
34 +Local:
35 +
36 +```bash
37 +pip install -r requirements.txt
38 +python -m countrysense.clean --raw-dir data/raw
39 +python -m countrysense.train --config configs/default.yaml
40 +python -m countrysense.evaluate --checkpoint runs/clip_vit_b16/best.pt
41 +python -m countrysense.predict --image path/to/photo.jpg
42 +```
43 +
44 +Optional Gradio demo: `pip install gradio` and `python app.py`.
45 +
46 +## Evaluation
47 +
48 +`evaluate.py` reports top-1, top-5 and macro F1, writes per-country accuracy and a confusion matrix. Top-5 matters for geography: confusing Estonia with Latvia is a much smaller error than confusing it with Brazil, and neighbor confusion dominates the mistakes.
49 +
50 +## Layout
51 +
52 +```
53 +countrysense/ package: data, clean, model, train, evaluate, predict
54 +configs/default.yaml training configuration
55 +notebooks/ Colab notebook, end to end
56 +docs/mudel.html architecture explainer (Estonian)
57 +app.py Gradio demo
58 +```
59 +
60 +## Roadmap
61 +
62 +- [ ] Region-level prediction within a country (hierarchical head: country, then region)
63 +- [ ] Larger training set (OpenStreetView-5M subset)
64 +- [ ] Calibration and abstention ("not confident enough to guess")
65 +- Not planned: exact coordinates. That is research territory (see PIGEON, GeoCLIP) and needs compute this project intentionally avoids.
added app.py +35 −0
@@ -0,0 +1,35 @@
1 +import os
2 +
3 +import gradio as gr
4 +import timm
5 +import torch
6 +
7 +from countrysense.data import build_transforms
8 +from countrysense.model import load_checkpoint
9 +
10 +CHECKPOINT = os.environ.get("COUNTRYSENSE_CKPT", "runs/clip_vit_b16/best.pt")
11 +
12 +device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
13 +model, classes = load_checkpoint(CHECKPOINT, device)
14 +data_cfg = timm.data.resolve_model_data_config(model.backbone)
15 +tfm = build_transforms(data_cfg["mean"], data_cfg["std"],
16 + data_cfg["input_size"][-1], train=False)
17 +
18 +
19 +def predict(img):
20 + with torch.no_grad():
21 + probs = model(tfm(img.convert("RGB")).unsqueeze(0).to(device)).softmax(1)[0]
22 + return {classes[i]: float(p) for p, i in
23 + zip(*[t.tolist() for t in probs.topk(min(5, len(classes)))])}
24 +
25 +
26 +demo = gr.Interface(
27 + fn=predict,
28 + inputs=gr.Image(type="pil"),
29 + outputs=gr.Label(num_top_classes=5),
30 + title="CountrySense",
31 + description="Ennustab tänavapildi järgi riigi. Riigi tase, mitte koordinaadid.",
32 +)
33 +
34 +if __name__ == "__main__":
35 + demo.launch()
added configs/default.yaml +24 −0
@@ -0,0 +1,24 @@
1 +data:
2 + raw_dir: data/raw # iga alamkaust on üks riik
3 + manifest: data/manifest.csv # clean.py väljund; kui puudub, skannitakse raw_dir
4 + min_per_class: 100 # riigid, millel on vähem pilte, jäetakse välja
5 + val_frac: 0.1
6 + test_frac: 0.1
7 + num_workers: 2
8 +
9 +model:
10 + backbone: vit_base_patch16_clip_224.openai
11 + dropout: 0.2
12 +
13 +train:
14 + seed: 42
15 + batch_size: 64
16 + epochs_head: 3 # faas 1: ainult klassifitseerimispea
17 + epochs_finetune: 8 # faas 2: viimased plokid lahti
18 + unfreeze_blocks: 4
19 + lr_head: 1.0e-3
20 + lr_head_finetune: 1.0e-4
21 + lr_backbone: 2.0e-5
22 + weight_decay: 0.01
23 + label_smoothing: 0.1
24 + out_dir: runs/clip_vit_b16
added countrysense/__init__.py +0 −0

Line changes are not available for this file.

added countrysense/clean.py +105 −0
@@ -0,0 +1,105 @@
1 +import argparse
2 +from pathlib import Path
3 +
4 +import imagehash
5 +import numpy as np
6 +import pandas as pd
7 +from PIL import Image
8 +from tqdm import tqdm
9 +
10 +from .data import scan_image_folder
11 +
12 +
13 +def sharpness_score(gray):
14 + # Gradiendi energia keskmine: udusel pildil on see madal.
15 + gy, gx = np.gradient(gray)
16 + return float(np.mean(gx * gx + gy * gy))
17 +
18 +
19 +def inspect_image(path, min_side, dark_limit, bright_limit, min_sharpness):
20 + # Tagastab (phash, None) kui pilt kõlbab, muidu (None, põhjus).
21 + try:
22 + with Image.open(path) as im:
23 + im = im.convert("RGB")
24 + except Exception:
25 + return None, "rikutud fail"
26 + if min(im.size) < min_side:
27 + return None, "liiga väike"
28 + small = np.asarray(im.convert("L").resize((128, 128)), dtype=np.float32)
29 + mean = small.mean()
30 + if mean < dark_limit:
31 + return None, "liiga tume"
32 + if mean > bright_limit:
33 + return None, "liiga hele"
34 + if sharpness_score(small) < min_sharpness:
35 + return None, "liiga udune"
36 + return imagehash.phash(im), None
37 +
38 +
39 +def main():
40 + ap = argparse.ArgumentParser(description="Puhastab toorandmestiku ja kirjutab manifesti")
41 + ap.add_argument("--raw-dir", required=True, help="kaust, kus iga alamkaust on üks riik")
42 + ap.add_argument("--out-manifest", default="data/manifest.csv")
43 + ap.add_argument("--report", default="data/cleaning_report.csv")
44 + ap.add_argument("--min-per-class", type=int, default=100)
45 + ap.add_argument("--max-per-class", type=int, default=2000)
46 + ap.add_argument("--min-side", type=int, default=128)
47 + ap.add_argument("--dark-limit", type=float, default=20.0)
48 + ap.add_argument("--bright-limit", type=float, default=235.0)
49 + ap.add_argument("--min-sharpness", type=float, default=25.0)
50 + ap.add_argument("--seed", type=int, default=42)
51 + args = ap.parse_args()
52 +
53 + df = scan_image_folder(args.raw_dir)
54 + print(f"Leitud {len(df)} pilti, {df['country'].nunique()} riiki")
55 +
56 + kept, dropped = [], []
57 + seen_hashes = {}
58 + for row in tqdm(df.itertuples(), total=len(df), desc="Kontrollin pilte"):
59 + phash, reason = inspect_image(
60 + row.path, args.min_side, args.dark_limit, args.bright_limit, args.min_sharpness
61 + )
62 + if reason is None:
63 + key = str(phash)
64 + if key in seen_hashes:
65 + reason = f"duplikaat ({seen_hashes[key]})"
66 + else:
67 + seen_hashes[key] = row.path
68 + kept.append({"path": row.path, "country": row.country})
69 + if reason is not None:
70 + dropped.append({"path": row.path, "country": row.country, "reason": reason})
71 +
72 + kept_df = pd.DataFrame(kept)
73 +
74 + # Klasside piirid: liiga väikesed riigid välja, liiga suured lakke.
75 + counts = kept_df["country"].value_counts()
76 + small = counts[counts < args.min_per_class].index
77 + for country in small:
78 + rows = kept_df[kept_df["country"] == country]
79 + dropped.extend(
80 + {"path": p, "country": country, "reason": "riigil liiga vähe pilte"}
81 + for p in rows["path"]
82 + )
83 + kept_df = kept_df[~kept_df["country"].isin(small)]
84 + capped = [
85 + g.sample(min(len(g), args.max_per_class), random_state=args.seed)
86 + for _, g in kept_df.groupby("country")
87 + ]
88 + kept_df = pd.concat(capped).reset_index(drop=True)
89 +
90 + Path(args.out_manifest).parent.mkdir(parents=True, exist_ok=True)
91 + kept_df.to_csv(args.out_manifest, index=False)
92 + Path(args.report).parent.mkdir(parents=True, exist_ok=True)
93 + dropped_df = pd.DataFrame(dropped, columns=["path", "country", "reason"])
94 + dropped_df.to_csv(args.report, index=False)
95 +
96 + print(f"Alles {len(kept_df)} pilti, {kept_df['country'].nunique()} riiki")
97 + if len(dropped_df):
98 + # Duplikaatide põhjus sisaldab faili teed, koondame üldnime alla.
99 + reasons = dropped_df["reason"].str.replace(r"duplikaat \(.*\)", "duplikaat", regex=True)
100 + print("Välja jäetud põhjuste kaupa:")
101 + print(reasons.value_counts().to_string())
102 +
103 +
104 +if __name__ == "__main__":
105 + main()
added countrysense/data.py +66 −0
@@ -0,0 +1,66 @@
1 +from pathlib import Path
2 +
3 +import pandas as pd
4 +from PIL import Image
5 +from torch.utils.data import Dataset
6 +from torchvision import transforms
7 +
8 +IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
9 +
10 +
11 +def scan_image_folder(root):
12 + # Ootab paigutust root/<riik>/<pilt>. Alamkausta nimi on klassi silt.
13 + root = Path(root)
14 + rows = []
15 + for country_dir in sorted(p for p in root.iterdir() if p.is_dir()):
16 + for img in sorted(country_dir.rglob("*")):
17 + if img.suffix.lower() in IMAGE_EXTS:
18 + rows.append({"path": str(img), "country": country_dir.name})
19 + if not rows:
20 + raise FileNotFoundError(f"Ei leidnud ühtegi pilti kaustast {root}")
21 + return pd.DataFrame(rows)
22 +
23 +
24 +def stratified_split(df, val_frac=0.1, test_frac=0.1, seed=42):
25 + # Jagame iga riigi sees eraldi, et val/test sisaldaks kõiki klasse.
26 + parts = {"train": [], "val": [], "test": []}
27 + for _, group in df.groupby("country"):
28 + 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))
31 + parts["val"].append(group.iloc[:n_val])
32 + parts["test"].append(group.iloc[n_val:n_val + n_test])
33 + parts["train"].append(group.iloc[n_val + n_test:])
34 + return {name: pd.concat(chunks).reset_index(drop=True) for name, chunks in parts.items()}
35 +
36 +
37 +def build_transforms(mean, std, image_size=224, train=True):
38 + if train:
39 + # Teadlikult ilma RandomHorizontalFlip'ita: liiklussuund (vasak- või
40 + # parempoolne) on päris geograafiline vihje, mille peegeldamine rikuks.
41 + return transforms.Compose([
42 + transforms.RandomResizedCrop(image_size, scale=(0.5, 1.0), ratio=(0.9, 1.1)),
43 + transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.02),
44 + transforms.ToTensor(),
45 + transforms.Normalize(mean, std),
46 + ])
47 + return transforms.Compose([
48 + transforms.Resize(image_size + 32),
49 + transforms.CenterCrop(image_size),
50 + transforms.ToTensor(),
51 + transforms.Normalize(mean, std),
52 + ])
53 +
54 +
55 +class CountryDataset(Dataset):
56 + def __init__(self, df, class_to_idx, transform):
57 + self.paths = df["path"].tolist()
58 + self.labels = [class_to_idx[c] for c in df["country"]]
59 + self.transform = transform
60 +
61 + def __len__(self):
62 + return len(self.paths)
63 +
64 + def __getitem__(self, idx):
65 + img = Image.open(self.paths[idx]).convert("RGB")
66 + return self.transform(img), self.labels[idx]
added countrysense/evaluate.py +99 −0
@@ -0,0 +1,99 @@
1 +import argparse
2 +import json
3 +from pathlib import Path
4 +
5 +import matplotlib
6 +matplotlib.use("Agg")
7 +import matplotlib.pyplot as plt
8 +import numpy as np
9 +import pandas as pd
10 +import timm
11 +import torch
12 +from sklearn.metrics import confusion_matrix, f1_score
13 +from torch.utils.data import DataLoader
14 +
15 +from .data import CountryDataset, build_transforms
16 +from .model import load_checkpoint
17 +
18 +
19 +@torch.no_grad()
20 +def collect_predictions(model, loader, device):
21 + preds, tops, targets = [], [], []
22 + for images, labels in loader:
23 + images = images.to(device, non_blocking=True)
24 + logits = model(images)
25 + k = min(5, logits.size(1))
26 + preds.append(logits.argmax(1).cpu())
27 + tops.append(logits.topk(k, dim=1).indices.cpu())
28 + targets.append(labels)
29 + return torch.cat(preds).numpy(), torch.cat(tops).numpy(), torch.cat(targets).numpy()
30 +
31 +
32 +def main():
33 + ap = argparse.ArgumentParser(description="Hindab mudelit testihulgal")
34 + ap.add_argument("--checkpoint", default="runs/clip_vit_b16/best.pt")
35 + ap.add_argument("--split-csv", default=None, help="vaikimisi test.csv checkpointi kaustast")
36 + ap.add_argument("--batch-size", type=int, default=64)
37 + ap.add_argument("--num-workers", type=int, default=2)
38 + ap.add_argument("--plot-top", type=int, default=30, help="mitu suurima toega riiki joonisele")
39 + args = ap.parse_args()
40 +
41 + ckpt_dir = Path(args.checkpoint).parent
42 + split_csv = Path(args.split_csv) if args.split_csv else ckpt_dir / "test.csv"
43 + device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
44 +
45 + model, classes = load_checkpoint(args.checkpoint, device)
46 + class_to_idx = {c: i for i, c in enumerate(classes)}
47 +
48 + df = pd.read_csv(split_csv)
49 + df = df[df["country"].isin(class_to_idx)].reset_index(drop=True)
50 + data_cfg = timm.data.resolve_model_data_config(model.backbone)
51 + ds = CountryDataset(
52 + df, class_to_idx,
53 + build_transforms(data_cfg["mean"], data_cfg["std"], data_cfg["input_size"][-1], train=False),
54 + )
55 + loader = DataLoader(ds, batch_size=args.batch_size, num_workers=args.num_workers)
56 +
57 + preds, tops, targets = collect_predictions(model, loader, device)
58 + top1 = float((preds == targets).mean())
59 + top5 = float((tops == targets[:, None]).any(1).mean())
60 + macro_f1 = float(f1_score(targets, preds, average="macro"))
61 +
62 + metrics = {"top1": top1, "top5": top5, "macro_f1": macro_f1,
63 + "n_test": int(len(targets)), "n_classes": len(classes)}
64 + with open(ckpt_dir / "metrics.json", "w", encoding="utf-8") as f:
65 + json.dump(metrics, f, indent=2)
66 + print(json.dumps(metrics, indent=2))
67 +
68 + per_country = (
69 + pd.DataFrame({"country": [classes[i] for i in targets], "correct": preds == targets})
70 + .groupby("country")
71 + .agg(accuracy=("correct", "mean"), support=("correct", "size"))
72 + .sort_values("accuracy")
73 + )
74 + per_country.to_csv(ckpt_dir / "per_country.csv")
75 + print("Nõrgimad riigid:")
76 + print(per_country.head(10).to_string())
77 +
78 + cm = confusion_matrix(targets, preds, labels=range(len(classes)))
79 + pd.DataFrame(cm, index=classes, columns=classes).to_csv(ckpt_dir / "confusion_matrix.csv")
80 +
81 + # Joonisele ainult suurima toega riigid, muidu on maatriks loetamatu.
82 + top_names = per_country.sort_values("support", ascending=False).head(args.plot_top).index
83 + idx = [class_to_idx[c] for c in top_names]
84 + sub = cm[np.ix_(idx, idx)].astype(np.float64)
85 + sub = sub / sub.sum(axis=1, keepdims=True).clip(min=1)
86 + fig, ax = plt.subplots(figsize=(12, 10))
87 + ax.imshow(sub, cmap="Blues", vmin=0, vmax=1)
88 + ax.set_xticks(range(len(idx)), top_names, rotation=90, fontsize=7)
89 + ax.set_yticks(range(len(idx)), top_names, fontsize=7)
90 + ax.set_xlabel("Ennustatud")
91 + ax.set_ylabel("Tegelik")
92 + ax.set_title(f"Confusion matrix, {args.plot_top} suurima toega riiki (reanormeeritud)")
93 + fig.tight_layout()
94 + fig.savefig(ckpt_dir / "confusion_matrix.png", dpi=150)
95 + print(f"Tulemused kaustas {ckpt_dir}")
96 +
97 +
98 +if __name__ == "__main__":
99 + main()
added countrysense/model.py +48 −0
@@ -0,0 +1,48 @@
1 +import timm
2 +import torch
3 +import torch.nn as nn
4 +
5 +
6 +class CountryClassifier(nn.Module):
7 + def __init__(self, num_classes, backbone="vit_base_patch16_clip_224.openai",
8 + dropout=0.2, pretrained=True):
9 + super().__init__()
10 + self.backbone = timm.create_model(backbone, pretrained=pretrained, num_classes=0)
11 + self.head = nn.Sequential(
12 + nn.Dropout(dropout),
13 + nn.Linear(self.backbone.num_features, num_classes),
14 + )
15 +
16 + def forward(self, x):
17 + return self.head(self.backbone(x))
18 +
19 + def freeze_backbone(self):
20 + for p in self.backbone.parameters():
21 + p.requires_grad = False
22 +
23 + def set_finetune_mode(self, unfreeze_blocks):
24 + self.freeze_backbone()
25 + if hasattr(self.backbone, "blocks"):
26 + for block in list(self.backbone.blocks)[-unfreeze_blocks:]:
27 + for p in block.parameters():
28 + p.requires_grad = True
29 + if hasattr(self.backbone, "norm"):
30 + for p in self.backbone.norm.parameters():
31 + p.requires_grad = True
32 + else:
33 + # Mitte-ViT backbone'il puudub plokkide loend, avame kõik.
34 + for p in self.backbone.parameters():
35 + p.requires_grad = True
36 +
37 +
38 +def load_checkpoint(path, device="cpu"):
39 + ckpt = torch.load(path, map_location=device, weights_only=True)
40 + model = CountryClassifier(
41 + num_classes=len(ckpt["classes"]),
42 + backbone=ckpt["backbone"],
43 + dropout=ckpt.get("dropout", 0.0),
44 + pretrained=False,
45 + )
46 + model.load_state_dict(ckpt["model"])
47 + model.to(device).eval()
48 + return model, ckpt["classes"]
added countrysense/predict.py +33 −0
@@ -0,0 +1,33 @@
1 +import argparse
2 +
3 +import timm
4 +import torch
5 +from PIL import Image
6 +
7 +from .data import build_transforms
8 +from .model import load_checkpoint
9 +
10 +
11 +def main():
12 + ap = argparse.ArgumentParser(description="Ennustab pildi riigi")
13 + ap.add_argument("--checkpoint", default="runs/clip_vit_b16/best.pt")
14 + ap.add_argument("--image", required=True)
15 + ap.add_argument("--topk", type=int, default=5)
16 + args = ap.parse_args()
17 +
18 + device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
19 + model, classes = load_checkpoint(args.checkpoint, device)
20 + data_cfg = timm.data.resolve_model_data_config(model.backbone)
21 + tfm = build_transforms(data_cfg["mean"], data_cfg["std"],
22 + data_cfg["input_size"][-1], train=False)
23 +
24 + img = Image.open(args.image).convert("RGB")
25 + with torch.no_grad():
26 + probs = model(tfm(img).unsqueeze(0).to(device)).softmax(1)[0]
27 + top = probs.topk(min(args.topk, len(classes)))
28 + for p, i in zip(top.values.tolist(), top.indices.tolist()):
29 + print(f"{classes[i]:<30s} {p:6.1%}")
30 +
31 +
32 +if __name__ == "__main__":
33 + main()
added countrysense/train.py +177 −0
@@ -0,0 +1,177 @@
1 +import argparse
2 +import csv
3 +import random
4 +from pathlib import Path
5 +
6 +import numpy as np
7 +import pandas as pd
8 +import timm
9 +import torch
10 +import torch.nn as nn
11 +import yaml
12 +from torch.utils.data import DataLoader, WeightedRandomSampler
13 +
14 +from .data import CountryDataset, build_transforms, scan_image_folder, stratified_split
15 +from .model import CountryClassifier
16 +
17 +
18 +def set_seed(seed):
19 + random.seed(seed)
20 + np.random.seed(seed)
21 + torch.manual_seed(seed)
22 + torch.cuda.manual_seed_all(seed)
23 +
24 +
25 +def train_one_epoch(model, loader, criterion, optimizer, scaler, device, use_amp):
26 + model.train()
27 + loss_sum, correct, total = 0.0, 0, 0
28 + for images, targets in loader:
29 + images = images.to(device, non_blocking=True)
30 + targets = targets.to(device, non_blocking=True)
31 + with torch.autocast(device_type=device.type, enabled=use_amp):
32 + logits = model(images)
33 + loss = criterion(logits, targets)
34 + optimizer.zero_grad(set_to_none=True)
35 + scaler.scale(loss).backward()
36 + scaler.unscale_(optimizer)
37 + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
38 + scaler.step(optimizer)
39 + scaler.update()
40 + loss_sum += loss.item() * targets.size(0)
41 + correct += (logits.argmax(1) == targets).sum().item()
42 + total += targets.size(0)
43 + return loss_sum / total, correct / total
44 +
45 +
46 +@torch.no_grad()
47 +def evaluate(model, loader, criterion, device, use_amp):
48 + model.eval()
49 + loss_sum, top1, top5, total = 0.0, 0, 0, 0
50 + for images, targets in loader:
51 + images = images.to(device, non_blocking=True)
52 + targets = targets.to(device, non_blocking=True)
53 + with torch.autocast(device_type=device.type, enabled=use_amp):
54 + logits = model(images)
55 + loss = criterion(logits, targets)
56 + loss_sum += loss.item() * targets.size(0)
57 + top1 += (logits.argmax(1) == targets).sum().item()
58 + k = min(5, logits.size(1))
59 + top5 += (logits.topk(k, dim=1).indices == targets.unsqueeze(1)).any(1).sum().item()
60 + total += targets.size(0)
61 + return loss_sum / total, top1 / total, top5 / total
62 +
63 +
64 +def main():
65 + ap = argparse.ArgumentParser(description="Treenib riigiklassifitseerija")
66 + ap.add_argument("--config", default="configs/default.yaml")
67 + args = ap.parse_args()
68 + with open(args.config, encoding="utf-8") as f:
69 + cfg = yaml.safe_load(f)
70 +
71 + d, m, t = cfg["data"], cfg["model"], cfg["train"]
72 + set_seed(t["seed"])
73 + device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
74 + use_amp = device.type == "cuda"
75 + out_dir = Path(t["out_dir"])
76 + out_dir.mkdir(parents=True, exist_ok=True)
77 +
78 + manifest = Path(d["manifest"])
79 + if manifest.exists():
80 + df = pd.read_csv(manifest)
81 + print(f"Manifest: {len(df)} pilti")
82 + else:
83 + print(f"Manifesti {manifest} pole, skannin {d['raw_dir']} (soovitav on enne käivitada clean.py)")
84 + df = scan_image_folder(d["raw_dir"])
85 + counts = df["country"].value_counts()
86 + df = df[df["country"].isin(counts[counts >= d["min_per_class"]].index)].reset_index(drop=True)
87 +
88 + classes = sorted(df["country"].unique())
89 + class_to_idx = {c: i for i, c in enumerate(classes)}
90 + print(f"{len(df)} pilti, {len(classes)} riiki, seade: {device}")
91 +
92 + splits = stratified_split(df, d["val_frac"], d["test_frac"], t["seed"])
93 + for name, part in splits.items():
94 + part.to_csv(out_dir / f"{name}.csv", index=False)
95 +
96 + model = CountryClassifier(len(classes), m["backbone"], m["dropout"]).to(device)
97 + data_cfg = timm.data.resolve_model_data_config(model.backbone)
98 + mean, std, image_size = data_cfg["mean"], data_cfg["std"], data_cfg["input_size"][-1]
99 +
100 + train_ds = CountryDataset(splits["train"], class_to_idx, build_transforms(mean, std, image_size, train=True))
101 + val_ds = CountryDataset(splits["val"], class_to_idx, build_transforms(mean, std, image_size, train=False))
102 +
103 + # Kaalutud valim: iga riik jõuab batch'idesse võrdse tõenäosusega,
104 + # muidu domineeriksid suurte piltide arvuga riigid.
105 + train_counts = splits["train"]["country"].value_counts()
106 + weights = splits["train"]["country"].map(lambda c: 1.0 / train_counts[c]).to_numpy()
107 + sampler = WeightedRandomSampler(
108 + torch.tensor(weights, dtype=torch.double),
109 + num_samples=len(weights),
110 + replacement=True,
111 + generator=torch.Generator().manual_seed(t["seed"]),
112 + )
113 + loader_kw = dict(batch_size=t["batch_size"], num_workers=d["num_workers"],
114 + pin_memory=use_amp, persistent_workers=d["num_workers"] > 0)
115 + train_loader = DataLoader(train_ds, sampler=sampler, drop_last=True, **loader_kw)
116 + val_loader = DataLoader(val_ds, shuffle=False, **loader_kw)
117 +
118 + criterion = nn.CrossEntropyLoss(label_smoothing=t["label_smoothing"])
119 + scaler = torch.amp.GradScaler(device.type, enabled=use_amp)
120 +
121 + log_path = out_dir / "log.csv"
122 + with open(log_path, "w", newline="", encoding="utf-8") as f:
123 + csv.writer(f).writerow(
124 + ["phase", "epoch", "lr", "train_loss", "train_acc", "val_loss", "val_top1", "val_top5"]
125 + )
126 +
127 + best_top1 = 0.0
128 +
129 + def log_and_save(phase, epoch, optimizer, train_loss, train_acc):
130 + nonlocal best_top1
131 + val_loss, val_top1, val_top5 = evaluate(model, val_loader, criterion, device, use_amp)
132 + lr = optimizer.param_groups[0]["lr"]
133 + with open(log_path, "a", newline="", encoding="utf-8") as f:
134 + csv.writer(f).writerow(
135 + [phase, epoch, f"{lr:.2e}", f"{train_loss:.4f}", f"{train_acc:.4f}",
136 + f"{val_loss:.4f}", f"{val_top1:.4f}", f"{val_top5:.4f}"]
137 + )
138 + print(f"[{phase}] epohh {epoch}: train_acc={train_acc:.3f} "
139 + f"val_top1={val_top1:.3f} val_top5={val_top5:.3f}")
140 + if val_top1 > best_top1:
141 + best_top1 = val_top1
142 + torch.save(
143 + {"model": model.state_dict(), "classes": classes,
144 + "backbone": m["backbone"], "dropout": m["dropout"], "val_top1": val_top1},
145 + out_dir / "best.pt",
146 + )
147 +
148 + # Faas 1: backbone külmutatud, treenime ainult klassifitseerimispead.
149 + # Suvaliselt initsialiseeritud pea gradiendid lõhuksid eeltreenitud kaale.
150 + model.freeze_backbone()
151 + optimizer = torch.optim.AdamW(model.head.parameters(), lr=t["lr_head"],
152 + weight_decay=t["weight_decay"])
153 + for epoch in range(1, t["epochs_head"] + 1):
154 + train_loss, train_acc = train_one_epoch(model, train_loader, criterion,
155 + optimizer, scaler, device, use_amp)
156 + log_and_save("head", epoch, optimizer, train_loss, train_acc)
157 +
158 + # Faas 2: viimased transformeri plokid lahti, madal LR ja koosinusgraafik.
159 + model.set_finetune_mode(t["unfreeze_blocks"])
160 + backbone_params = [p for p in model.backbone.parameters() if p.requires_grad]
161 + optimizer = torch.optim.AdamW(
162 + [{"params": backbone_params, "lr": t["lr_backbone"]},
163 + {"params": model.head.parameters(), "lr": t["lr_head_finetune"]}],
164 + weight_decay=t["weight_decay"],
165 + )
166 + scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=t["epochs_finetune"])
167 + for epoch in range(1, t["epochs_finetune"] + 1):
168 + train_loss, train_acc = train_one_epoch(model, train_loader, criterion,
169 + optimizer, scaler, device, use_amp)
170 + log_and_save("finetune", epoch, optimizer, train_loss, train_acc)
171 + scheduler.step()
172 +
173 + print(f"Valmis. Parim val top-1: {best_top1:.3f}, checkpoint: {out_dir / 'best.pt'}")
174 +
175 +
176 +if __name__ == "__main__":
177 + main()
added docs/mudel.html +403 −0
@@ -0,0 +1,403 @@
1 +<!DOCTYPE html>
2 +<html lang="et">
3 +<head>
4 +<meta charset="utf-8">
5 +<meta name="viewport" content="width=device-width, initial-scale=1">
6 +<title>CountrySense: kuidas mudel töötab</title>
7 +<style>
8 + :root {
9 + --ink: #1d2733;
10 + --muted: #5b6b7c;
11 + --line: #d9e1ea;
12 + --accent: #1a6feb;
13 + --accent-soft: #eaf2ff;
14 + --bg: #fbfcfe;
15 + --code-bg: #f2f5f9;
16 + }
17 + * { box-sizing: border-box; }
18 + body {
19 + margin: 0;
20 + font-family: Georgia, "Times New Roman", serif;
21 + color: var(--ink);
22 + background: var(--bg);
23 + line-height: 1.65;
24 + }
25 + header {
26 + background: var(--ink);
27 + color: #fff;
28 + padding: 3rem 1.5rem 2.5rem;
29 + text-align: center;
30 + }
31 + header h1 { margin: 0 0 0.4rem; font-size: 2.2rem; }
32 + header p { margin: 0; color: #b9c6d4; font-size: 1.05rem; }
33 + main { max-width: 860px; margin: 0 auto; padding: 1.5rem; }
34 + nav.toc {
35 + background: #fff;
36 + border: 1px solid var(--line);
37 + border-radius: 8px;
38 + padding: 1rem 1.5rem;
39 + margin: 1.5rem 0;
40 + font-family: system-ui, sans-serif;
41 + font-size: 0.92rem;
42 + }
43 + nav.toc ol { margin: 0.3rem 0 0; padding-left: 1.3rem; columns: 2; }
44 + nav.toc a { color: var(--accent); text-decoration: none; }
45 + nav.toc a:hover { text-decoration: underline; }
46 + h2 {
47 + margin-top: 2.6rem;
48 + padding-bottom: 0.3rem;
49 + border-bottom: 2px solid var(--line);
50 + font-size: 1.5rem;
51 + }
52 + h3 { margin-top: 1.8rem; font-size: 1.15rem; }
53 + code, pre {
54 + font-family: Consolas, "Courier New", monospace;
55 + background: var(--code-bg);
56 + border-radius: 4px;
57 + }
58 + code { padding: 0.1em 0.35em; font-size: 0.9em; }
59 + pre { padding: 0.9rem 1.1rem; overflow-x: auto; border: 1px solid var(--line); }
60 + pre code { background: none; padding: 0; }
61 + table {
62 + border-collapse: collapse;
63 + width: 100%;
64 + margin: 1rem 0;
65 + font-family: system-ui, sans-serif;
66 + font-size: 0.9rem;
67 + background: #fff;
68 + }
69 + th, td { border: 1px solid var(--line); padding: 0.5rem 0.7rem; text-align: left; vertical-align: top; }
70 + th { background: var(--accent-soft); }
71 + .note {
72 + background: var(--accent-soft);
73 + border-left: 4px solid var(--accent);
74 + padding: 0.8rem 1.1rem;
75 + border-radius: 0 6px 6px 0;
76 + margin: 1.2rem 0;
77 + }
78 + .warn {
79 + background: #fff6ec;
80 + border-left: 4px solid #e08a2e;
81 + padding: 0.8rem 1.1rem;
82 + border-radius: 0 6px 6px 0;
83 + margin: 1.2rem 0;
84 + }
85 + .pipeline {
86 + display: flex;
87 + flex-wrap: wrap;
88 + align-items: center;
89 + justify-content: center;
90 + gap: 0.4rem;
91 + margin: 1.5rem 0;
92 + font-family: system-ui, sans-serif;
93 + font-size: 0.82rem;
94 + }
95 + .pipeline .box {
96 + background: #fff;
97 + border: 1.5px solid var(--ink);
98 + border-radius: 6px;
99 + padding: 0.55rem 0.8rem;
100 + text-align: center;
101 + min-width: 105px;
102 + }
103 + .pipeline .box small { display: block; color: var(--muted); }
104 + .pipeline .arrow { font-size: 1.2rem; color: var(--muted); }
105 + .pipeline .box.hot { background: var(--accent-soft); border-color: var(--accent); }
106 + .block-diagram {
107 + background: #fff;
108 + border: 1px dashed var(--muted);
109 + border-radius: 8px;
110 + padding: 1rem 1.3rem;
111 + font-family: Consolas, monospace;
112 + font-size: 0.85rem;
113 + white-space: pre;
114 + overflow-x: auto;
115 + margin: 1.2rem 0;
116 + }
117 + .formula {
118 + text-align: center;
119 + font-size: 1.1rem;
120 + margin: 1.2rem 0;
121 + font-style: italic;
122 + }
123 + footer {
124 + margin-top: 3rem;
125 + padding: 1.5rem;
126 + text-align: center;
127 + color: var(--muted);
128 + font-size: 0.85rem;
129 + border-top: 1px solid var(--line);
130 + font-family: system-ui, sans-serif;
131 + }
132 +</style>
133 +</head>
134 +<body>
135 +
136 +<header>
137 + <h1>CountrySense</h1>
138 + <p>Kuidas nägemistransformer tänavapildi järgi riigi ära arvab: arhitektuur, andmed ja kõik valikud lahti seletatuna</p>
139 +</header>
140 +
141 +<main>
142 +
143 +<nav class="toc">
144 + <strong>Sisukord</strong>
145 + <ol>
146 + <li><a href="#skoop">Ülesanne ja skoop</a></li>
147 + <li><a href="#andmed">Andmestik ja valikukriteeriumid</a></li>
148 + <li><a href="#puhastus">Andmete puhastamine</a></li>
149 + <li><a href="#mudelivalik">Mudeli valik</a></li>
150 + <li><a href="#vit">ViT samm-sammult</a></li>
151 + <li><a href="#dense">Dense-kihid transformeri sees</a></li>
152 + <li><a href="#pea">Klassifitseerimispea ja softmax</a></li>
153 + <li><a href="#kadu">Kadufunktsioon</a></li>
154 + <li><a href="#treening">Treening kahes faasis</a></li>
155 + <li><a href="#augment">Augmentatsioonid</a></li>
156 + <li><a href="#tasakaal">Klasside tasakaal</a></li>
157 + <li><a href="#moodikud">Mõõdikud</a></li>
158 + <li><a href="#piirangud">Piirangud ja teekaart</a></li>
159 + </ol>
160 +</nav>
161 +
162 +<h2 id="skoop">1. Ülesanne ja skoop</h2>
163 +<p>
164 +Sisend on üks tänavapilt, väljund on riigi nimi ja tõenäosused. Kogu toru näeb välja nii:
165 +</p>
166 +
167 +<div class="pipeline">
168 + <div class="box">Pilt<small>224 × 224 × 3</small></div>
169 + <div class="arrow">→</div>
170 + <div class="box">Patch embedding<small>196 + 1 tokenit</small></div>
171 + <div class="arrow">→</div>
172 + <div class="box hot">12 × Transformeri plokk<small>attention + MLP</small></div>
173 + <div class="arrow">→</div>
174 + <div class="box">CLS-vektor<small>768 arvu</small></div>
175 + <div class="arrow">→</div>
176 + <div class="box hot">Dropout + Linear<small>768 → N riiki</small></div>
177 + <div class="arrow">→</div>
178 + <div class="box">Softmax<small>tõenäosused</small></div>
179 +</div>
180 +
181 +<p>
182 +Miks just riigi tase? Täpne geolokatsioon (koordinaatide ennustamine) on teadusartikli mõõtu ettevõtmine: miljonid pildid, suured mudelid, nädalad GPU-aega. Riigi klassifitseerimine saja klassi vahel on tavaline juhendatud õppe ülesanne, mis mahub tasuta Colabi GPU peale. Järgmine aus samm oleks regioon riigi sees, koordinaadid ei ole plaanis.
183 +</p>
184 +
185 +<h2 id="andmed">2. Andmestik ja valikukriteeriumid</h2>
186 +<p>
187 +Vaikimisi andmestik on Kaggle'i <em>GeoLocation, Geoguessr Images 50K</em>: umbes 50 000 Street View pilti, iga riik oma kaustas. Valik tehti nelja kriteeriumi järgi:
188 +</p>
189 +<table>
190 + <tr><th>Kriteerium</th><th>Miks oluline</th></tr>
191 + <tr><td>Sildid kaustastruktuuris</td><td>Käsitsi märgendamine oleks nädalate töö. Kausta nimi ongi klassi silt.</td></tr>
192 + <tr><td>~150 riiki, ~50k pilti</td><td>Piisavalt suur, et ülesanne oleks päris, ja piisavalt väike, et tasuta GPU jaksaks.</td></tr>
193 + <tr><td>Maht mõni gigabait</td><td>Mahub Colabi kettale, allalaadimine minutites, mitte tundides.</td></tr>
194 + <tr><td>Vaba juurdepääs</td><td>kagglehub laadib ilma API võtmeta, notebook töötab igaühel.</td></tr>
195 +</table>
196 +<p>
197 +Kood ei sõltu sellest konkreetsest andmestikust: iga kaust paigutusega <code>root/&lt;riik&gt;/&lt;pilt&gt;</code> sobib. Skaleerimiseks on hea kandidaat OpenStreetView-5M alamhulk.
198 +</p>
199 +<div class="warn">
200 +<strong>Ausalt biasest.</strong> Street View katvus on kaldu jõukamate riikide poole ja eri riikides on pildistatud eri kaamerapõlvkondadega. Mudel õpib paratamatult ka kaamera artefakte (värvitoon, teravus, resolutsioon), mitte ainult maastikku ja arhitektuuri. GeoGuessri profimängijad kasutavad täpselt sama vihjet.
201 +</div>
202 +
203 +<h2 id="puhastus">3. Andmete puhastamine</h2>
204 +<p>
205 +Toorandmestikus on katkiseid faile, praktiliselt musti kaadreid (tunnelid, öö), uduseid pilte ja duplikaate. <code>clean.py</code> käib kõik pildid läbi ja kirjutab iga välja jäetud pildi kohta raporti reaga põhjuse.
206 +</p>
207 +<table>
208 + <tr><th>Samm</th><th>Reegel</th><th>Miks</th></tr>
209 + <tr><td>Rikutud failid</td><td>PIL ei suuda avada</td><td>Katkine fail kukutaks treeningu keset epohhi.</td></tr>
210 + <tr><td>Suurus</td><td>lühem külg &lt; 128 px</td><td>Alla selle pole piisavalt detaili, ViT sisend on 224 px.</td></tr>
211 + <tr><td>Heledus</td><td>halltoonide keskmine &lt; 20 või &gt; 235</td><td>Peaaegu mustad või läbipõlenud kaadrid ei kanna infot, aga kannavad silti, st ainult müra.</td></tr>
212 + <tr><td>Teravus</td><td>gradiendi energia keskmine &lt; 25</td><td>Udune kaader (liikumine, vihmapiisk objektiivil). Gradiendi energia on lihtne teravuse mõõt: udusel pildil muutuvad naaberpikslid vähe.</td></tr>
213 + <tr><td>Duplikaadid</td><td>korduv phash</td><td>Perceptual hash annab visuaalselt identsetele piltidele sama koodi. Kriitiline: kui sama koht satub nii treening- kui testihulka, on testitulemus petlikult hea. See on andmeleke, kõige levinum viga sedasorti projektides.</td></tr>
214 + <tr><td>Väikesed klassid</td><td>riigil &lt; 100 pilti</td><td>Paarikümne pildiga ei saa õppida ega usaldusväärselt mõõta.</td></tr>
215 + <tr><td>Suured klassid</td><td>lagi 2000 pilti riigi kohta</td><td>Vähendab tasakaalustamatust juba andmete tasandil, ülejäänu teeb sampler (ptk 11).</td></tr>
216 +</table>
217 +
218 +<h2 id="mudelivalik">4. Mudeli valik</h2>
219 +<p>
220 +Otsustuskriteeriumid: eeltreeningu kvaliteet geograafia jaoks, parameetrite maht (peab mahtuma tasuta T4 GPU 16 GB mällu koos peenhäälestusega) ja teegi tugi.
221 +</p>
222 +<table>
223 + <tr><th>Kandidaat</th><th>Hinnang</th></tr>
224 + <tr><td><strong>CLIP ViT-B/16 (valitud)</strong></td><td>OpenAI treenis seda 400 miljoni pilt-tekst paari peal. Kuna tekstid kirjeldasid ka kohti ("a street in Lisbon"), kodeerivad tunnused juba silte, taimestikku, arhitektuuri, teekattemärgistust. Lineaarne pea CLIP-i tunnuste otsas on teadaolevalt tugev geolokatsiooni baastase. 86M parameetrit mahub T4 peale.</td></tr>
225 + <tr><td>ImageNet ViT / ConvNeXt</td><td>Töötaks, aga ImageNeti 1000 klassi (koeratõud, esemed) on geograafiast kaugel, transfer on mõõdetavalt nõrgem.</td></tr>
226 + <tr><td>StreetCLIP (ViT-L/14)</td><td>Juba geolokatsiooniks häälestatud CLIP, aga kolm korda suurem. Peenhäälestus T4 peal on kitsas. Hea järgmine samm parema riistvaraga.</td></tr>
227 + <tr><td>CNN nullist</td><td>50k pildiga lootusetu. Eeltreening on kogu projekti võimaldaja: keegi teine on juba kulutanud tuhanded GPU-tunnid üldiste visuaalsete tunnuste õppimisele.</td></tr>
228 +</table>
229 +
230 +<h3>ViT-B/16 numbrid</h3>
231 +<table>
232 + <tr><th>Omadus</th><th>Väärtus</th></tr>
233 + <tr><td>Parameetreid</td><td>~86 miljonit</td></tr>
234 + <tr><td>Transformeri plokke</td><td>12</td></tr>
235 + <tr><td>Peidetud dimensioon</td><td>768</td></tr>
236 + <tr><td>Attention-päid ploki kohta</td><td>12</td></tr>
237 + <tr><td>MLP vahekihi laius</td><td>3072</td></tr>
238 + <tr><td>Patchi suurus</td><td>16 × 16 px</td></tr>
239 + <tr><td>Tokeneid 224 px pildi kohta</td><td>196 patchi + 1 CLS = 197</td></tr>
240 +</table>
241 +
242 +<h2 id="vit">5. ViT samm-sammult</h2>
243 +
244 +<h3>5.1 Patch embedding: pilt muutub jadaks</h3>
245 +<p>
246 +Transformer töötab tokenite jadaga, mitte pikslivõrega. Seepärast lõigatakse 224 × 224 pilt 16 × 16 piksliseks ruudustikuks: 14 × 14 = 196 patchi. Iga patch on 16 × 16 × 3 = 768 arvu, mis lastakse läbi <em>ühe lineaarse kihi</em> (see on esimene dense-kiht kogu mudelis) ja saadakse 768-mõõtmeline vektor. Pilt on nüüd 196 "sõnast" koosnev lause.
247 +</p>
248 +<h3>5.2 CLS-token ja positsioonivektorid</h3>
249 +<p>
250 +Jada ette lisatakse üks õpitav lisatoken, <strong>CLS</strong> (classification). Tal pole pildisisu, tema ülesanne on plokkide läbimise käigus koguda attention'i kaudu kokku terve pildi kokkuvõte. Lõpus loetakse ennustus just tema pealt.
251 +</p>
252 +<p>
253 +Kuna attention ise on järjekorra suhtes ükskõikne, liidetakse igale tokenile <strong>positsioonivektor</strong>: õpitav vektor, mis ütleb "sina oled rea 3, veeru 7 patch". Ilma selleta ei teaks mudel, kas taevas on üleval või all.
254 +</p>
255 +
256 +<h3>5.3 Transformeri plokk</h3>
257 +<p>Kõik 12 plokki on identse ehitusega:</p>
258 +<div class="block-diagram">sisend (197 × 768)
259 + │
260 + ├──────────────────────────────┐
261 + ▼ │
262 +LayerNorm │
263 + ▼ │
264 +Multi-head self-attention │ (12 pead)
265 + ▼ │
266 + + ◄───────────────────────────┘ residuaalühendus
267 + │
268 + ├──────────────────────────────┐
269 + ▼ │
270 +LayerNorm │
271 + ▼ │
272 +MLP: Linear 768 → 3072 │ (dense-kihid)
273 + GELU │
274 + Linear 3072 → 768 │
275 + ▼ │
276 + + ◄───────────────────────────┘ residuaalühendus
277 + │
278 +väljund (197 × 768)</div>
279 +
280 +<h3>5.4 Self-attention: info liigub patchide vahel</h3>
281 +<p>
282 +Igast tokenist arvutatakse kolm projektsiooni: päring <strong>Q</strong> (query), võti <strong>K</strong> (key) ja väärtus <strong>V</strong> (value). Iga token "küsib" kõigilt teistelt, kui asjakohased nad talle on, ja korjab nende väärtused kokku kaalutult:
283 +</p>
284 +<p class="formula">Attention(Q, K, V) = softmax(Q·K<sup>T</sup> / √d) · V</p>
285 +<p>
286 +Jagamine √d-ga hoiab korrutised mõistlikus vahemikus, et softmax ei küllastuks. "Multi-head" tähendab, et see arvutus tehakse 12 korda paralleelselt väiksemates alamruumides (768 / 12 = 64 mõõdet pea kohta): üks pea võib jälgida värve, teine geomeetriat, kolmas teksti silmapiiril. Just attention lubab CLS-tokenil siduda kokku liiklusmärgi paremal, taimestiku vasakul ja teekattemärgistuse all, ükskõik kui kaugel need üksteisest on. CNN-il kuluks sama kaugete seoste jaoks palju kihte.
287 +</p>
288 +
289 +<h2 id="dense">6. Dense-kihid transformeri sees</h2>
290 +<p>
291 +Iga ploki teine pool on MLP (multi-layer perceptron), kaks dense- ehk täisühendatud kihti:
292 +</p>
293 +<pre><code>Linear(768 → 3072) # laiendus 4×
294 +GELU # sujuv mittelineaarsus
295 +Linear(3072 → 768) # tagasi kokku</code></pre>
296 +<p>
297 +Tööjaotus ploki sees on selge: <strong>attention liigutab infot tokenite vahel, MLP töötleb iga tokenit eraldi</strong>. MLP-s toimub tegelik tunnuste teisendamine: laiendus 3072 mõõtmesse annab ruumi vahepealsete kombinatsioonide jaoks ("kollane + ristkülik + posti otsas"), GELU teeb teisenduse mittelineaarseks (ilma selleta oleks kogu võrk üks suur maatrikskorrutis) ja tagasiprojektsioon surub tulemuse standardsesse 768-mõõtmelisse esitusse, mida järgmine plokk ootab.
298 +</p>
299 +<p>
300 +Mahult on MLP-d mudeli põhiosa: umbes kaks kolmandikku kõigist parameetritest elab just neis dense-kihtides.
301 +</p>
302 +<p>
303 +Kaks tugistruktuuri teevad 12 ploki virna treenitavaks: <strong>LayerNorm</strong> normaliseerib iga tokeni vektori enne igat alamosa (stabiilsed suurusjärgud), <strong>residuaalühendused</strong> (x + f(x)) lasevad gradiendil voolata otse läbi kogu virna, nii et sügav võrk ei "unusta" sisendit ega lämmata gradienti.
304 +</p>
305 +
306 +<h2 id="pea">7. Klassifitseerimispea ja softmax</h2>
307 +<p>
308 +Pärast 12. plokki ja lõpu LayerNorm'i võetakse CLS-tokeni 768-mõõtmeline vektor. See on kogu pildi kokkuvõte. Pea on tahtlikult minimaalne:
309 +</p>
310 +<pre><code>nn.Sequential(
311 + nn.Dropout(0.2), # treeningul nullitakse juhuslikult 20% tunnustest
312 + nn.Linear(768, N_riiki), # üks dense-kiht: 768 sisendit → N logitit
313 +)</code></pre>
314 +<p>
315 +<strong>Dropout</strong> takistab peal toetumast üksikutele tunnustele (näiteks ainult ühe kaamerapõlvkonna värvitoonile): kui iga tunnus võib treeningul kaduda, peab otsus toetuma laiemale mustrile.
316 +</p>
317 +<p>
318 +<strong>Linear</strong> annab igale riigile ühe reaalarvu, logiti. Sisuliselt on igal riigil 768 kaalu, mis ütlevad, millised tunnusekombinatsioonid tema poolt räägivad. <strong>Softmax</strong> teisendab logitid tõenäosusteks: e astmes iga logit, jagatud summaga, nii et tulemused on positiivsed ja annavad kokku 1.
319 +</p>
320 +<div class="note">
321 +<strong>Miks ainult üks kiht, mitte sügavam pea?</strong> CLIP-i tunnused on juba suures osas lineaarselt eraldatavad, seda mõõdab nn linear probe. Sügavam pea lisaks selle andmemahu (~40k treeningpilti) juures peamiselt ülesobitumise riski. Kui tunnused vajavad painutamist, on parem avada backbone'i viimased plokid, mida faasis 2 teemegi.
322 +</div>
323 +
324 +<h2 id="kadu">8. Kadufunktsioon</h2>
325 +<p>
326 +Cross-entropy võrdleb softmax'i tõenäosusjaotust õige vastusega ja karistab logaritmiliselt: kindel vale vastus on väga kallis. Lisaks kasutame <strong>label smoothing</strong> väärtusega 0.1: sihtmärk pole "Eesti = 1.0, kõik muu = 0.0", vaid "Eesti = 0.9, ülejäänu jagab 0.1".
327 +</p>
328 +<p>
329 +Põhjus on geograafiline: sildid on piirialadel loomupäraselt mürarikkad. Pilt Valgast näeb välja nagu Valka, ja Street View asukohaviga võib panna pildi sõna otseses mõttes vale riigi kausta. Smoothing hoiab mudelit sellistel juhtudel ülekindlaks muutumast ja parandab kalibratsiooni.
330 +</p>
331 +
332 +<h2 id="treening">9. Treening kahes faasis</h2>
333 +<table>
334 + <tr><th></th><th>Faas 1: lineaarne sondeerimine</th><th>Faas 2: osaline peenhäälestus</th></tr>
335 + <tr><td>Mis treenib</td><td>ainult pea (Dropout + Linear)</td><td>viimased 4 plokki + lõpu norm + pea</td></tr>
336 + <tr><td>Epohhe</td><td>3</td><td>8</td></tr>
337 + <tr><td>Õppemäär</td><td>1e-3</td><td>2e-5 backbone, 1e-4 pea, koosinusgraafik</td></tr>
338 +</table>
339 +<p>
340 +<strong>Miks mitte kohe kõike treenida?</strong> Pea alustab juhuslikest kaaludest ja tema esimesed gradiendid on suured ja suvalises suunas. Kui backbone oleks lahti, lammutaksid need gradiendid eeltreenitud kaale enne, kui pea üldse midagi mõistlikku nõuab. Seepärast laseme peal esmalt külmutatud tunnuste otsas paika loksuda.
341 +</p>
342 +<p>
343 +<strong>Miks ainult viimased 4 plokki, mitte kõik 12?</strong> Alumised plokid õpivad üldisi asju (servad, tekstuurid, värvilaigud), mis on igasuguse pildiülesande jaoks niigi head. Ülemised plokid kannavad semantikat, mida tasub geograafiale kohandada. Vähem avatud plokke tähendab kolme asja korraga: väiksem mälukulu (mahume T4 peale), väiksem ülesobitumise risk ja väiksem katastroofilise unustamise risk (et mudel ei kaotaks CLIP-i üldteadmisi).
344 +</p>
345 +<p>
346 +Tehnilised valikud: AdamW (weight decay 0.01 hoiab kaalud väiksed), AMP ehk poolttäpsusega arvutus (ligi 2× kiirem ja poole väiksem mälukulu), gradientide kärpimine normini 1.0 (üksik halb batch ei löö treeningut rööpast), koosinusgraafik faasis 2 (õppemäär langeb sujuvalt, lõpus tehakse peeneid samme).
347 +</p>
348 +
349 +<h2 id="augment">10. Augmentatsioonid</h2>
350 +<table>
351 + <tr><th>Teisendus</th><th>Kasutusel?</th><th>Põhjendus</th></tr>
352 + <tr><td>RandomResizedCrop (50 kuni 100% pindalast)</td><td>Jah</td><td>Kaugus ja kadreering varieeruvad päriselt: sama koht näeb eri suumiga erinev välja.</td></tr>
353 + <tr><td>ColorJitter (kerge)</td><td>Jah</td><td>Valgustus, aastaaeg ja kaamera värvitoon varieeruvad päriselt.</td></tr>
354 + <tr><td>Horisontaalne peegeldus</td><td><strong>Ei</strong></td><td>Vasak- või parempoolne liiklus on päris geograafiline tunnus. Peegeldatud Suurbritannia näeb välja nagu Prantsusmaa, silt jääb "Suurbritannia": õpetaksime mudelile otseselt valet.</td></tr>
355 + <tr><td>Pööramine</td><td>Ei</td><td>Horisont on tänavapildil alati enam-vähem loodis, kaldus pilte päriselt ei tule. Horisondi asend on info.</td></tr>
356 +</table>
357 +<div class="note">
358 +Hea augmentatsiooni reegel: simuleeri variatsiooni, mis päris andmetes olemas on, ja mitte kunagi sellist, mis silti muudaks.
359 +</div>
360 +
361 +<h2 id="tasakaal">11. Klasside tasakaal</h2>
362 +<p>
363 +Pärast puhastust on riikide vahel ikkagi suurusjärguline vahe: mõnel riigil 2000 pilti, mõnel 100. Ilma sekkumiseta õpiks mudel "kahtluse korral paku suurt riiki" ja väikesed riigid jääksid nulli.
364 +</p>
365 +<p>
366 +Lahendus on <code>WeightedRandomSampler</code>: iga pildi valikutõenäosus on pöördvõrdeline tema riigi piltide arvuga, nii et iga batch sisaldab kõiki riike ligikaudu võrdselt. Väikeste riikide pilte näidatakse lihtsalt sagedamini (tagasipanekuga valik). Alternatiiv oleks klassikaalud kadufunktsioonis; sampler annab ühtlasemad batchid ja stabiilsema treeningu, seepärast valisime tema.
367 +</p>
368 +
369 +<h2 id="moodikud">12. Mõõdikud</h2>
370 +<table>
371 + <tr><th>Mõõdik</th><th>Mida mõõdab</th><th>Miks siin oluline</th></tr>
372 + <tr><td>Top-1</td><td>esimene pakkumine õige</td><td>Põhinäitaja. Juhuslik pakkumine ~100 riigiga oleks 1%.</td></tr>
373 + <tr><td>Top-5</td><td>õige riik viie parima seas</td><td>Geograafias aus: Eesti-Läti segiajamine on väike viga, Eesti-Brasiilia suur. Vigade seas domineerivad naabrid.</td></tr>
374 + <tr><td>Makro-F1</td><td>klasside keskmine F1, kõik riigid võrdse kaaluga</td><td>Paljastab, kui mudel ohverdab väikesed riigid. Top-1 üksi seda ei näita.</td></tr>
375 + <tr><td>Riigipõhine tabel</td><td>täpsus ja tugi riigi kaupa</td><td>Näitab, kus mudel on nõrk ja kuhu andmeid juurde vaja.</td></tr>
376 + <tr><td>Confusion matrix</td><td>mida millega segi aetakse</td><td>Tüüpilised klastrid: Balti riigid, Skandinaavia, hispaaniakeelne Ladina-Ameerika. Kui segadused on geograafiliselt loogilised, õpib mudel õigeid tunnuseid.</td></tr>
377 +</table>
378 +<p>
379 +Testihulk on treeningust rangelt lahus (stratifitseeritud 80/10/10 jaotus, iga riigi sees eraldi) ja duplikaadid eemaldati enne jaotust, et sama koht ei satuks mõlemale poole.
380 +</p>
381 +
382 +<h2 id="piirangud">13. Piirangud ja teekaart</h2>
383 +<p>
384 +Mudel õpib osalt kaamerat, mitte maastikku (Street View põlvkondade bias). Riigid, mida andmestikus pole, saavad paratamatult mõne naabri sildi. Sisepildid ja inimesed pole selle mudeli teema, ta on tänavavaadete klassifitseerija.
385 +</p>
386 +<p>
387 +Teekaart, ausas järjekorras:
388 +</p>
389 +<ol>
390 + <li>Regioonitase riigi sees: hierarhiline pea, mis ennustab enne riigi ja siis regiooni. Sama retsept, rohkem klasse, rohkem andmeid.</li>
391 + <li>Suurem andmestik (OpenStreetView-5M alamhulk) ja suurem backbone (StreetCLIP), kui riistvara lubab.</li>
392 + <li>Kalibratsioon ja keeldumine: kui top-1 tõenäosus on madal, on ausam öelda "ei tea".</li>
393 + <li><strong>Mitte plaanis:</strong> täpsed koordinaadid. See on teadustöö territoorium (PIGEON, GeoCLIP) ja nõuab arvutusmahtu, mida see projekt teadlikult väldib.</li>
394 +</ol>
395 +
396 +</main>
397 +
398 +<footer>
399 +CountrySense · CLIP ViT-B/16 peenhäälestus riigi klassifitseerimiseks · vt ka README.et.md
400 +</footer>
401 +
402 +</body>
403 +</html>
added notebooks/train_colab.ipynb +184 −0
@@ -0,0 +1,184 @@
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 + "cells": [
19 + {
20 + "cell_type": "markdown",
21 + "metadata": {},
22 + "source": [
23 + "# CountrySense: treening Google Colabis\n",
24 + "\n",
25 + "Terve toru: andmestiku allalaadimine, puhastamine, treening kahes faasis ja hindamine.\n",
26 + "\n",
27 + "Enne alustamist: Runtime > Change runtime type > T4 GPU.\n",
28 + "\n",
29 + "Kogu treening kestab T4 peal suurusjärgus 1 kuni 2 tundi."
30 + ]
31 + },
32 + {
33 + "cell_type": "code",
34 + "metadata": {},
35 + "execution_count": null,
36 + "outputs": [],
37 + "source": [
38 + "!nvidia-smi"
39 + ]
40 + },
41 + {
42 + "cell_type": "code",
43 + "metadata": {},
44 + "execution_count": null,
45 + "outputs": [],
46 + "source": [
47 + "# Asenda URL oma repoga.\n",
48 + "!git clone https://github.com/rasmju/countrysense.git\n",
49 + "%cd countrysense\n",
50 + "!pip install -q -r requirements.txt kagglehub"
51 + ]
52 + },
53 + {
54 + "cell_type": "markdown",
55 + "metadata": {},
56 + "source": [
57 + "## Andmestik\n",
58 + "\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."
60 + ]
61 + },
62 + {
63 + "cell_type": "code",
64 + "metadata": {},
65 + "execution_count": null,
66 + "outputs": [],
67 + "source": [
68 + "import kagglehub\n",
69 + "from pathlib import Path\n",
70 + "\n",
71 + "ds_path = Path(kagglehub.dataset_download(\"ubitquitin/geolocation-geoguessr-images-50k\"))\n",
72 + "\n",
73 + "# Leiame kausta, mille alamkaustad on riigid (paigutus voib versiooniti erineda).\n",
74 + "candidates = [p for p in [ds_path, *ds_path.rglob(\"*\")] if p.is_dir()]\n",
75 + "raw_dir = max(candidates, key=lambda p: sum(1 for c in p.iterdir() if c.is_dir()))\n",
76 + "print(\"Toorandmed:\", raw_dir)\n",
77 + "print(\"Riikide kaustu:\", sum(1 for c in raw_dir.iterdir() if c.is_dir()))"
78 + ]
79 + },
80 + {
81 + "cell_type": "markdown",
82 + "metadata": {},
83 + "source": [
84 + "## Puhastamine\n",
85 + "\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`."
87 + ]
88 + },
89 + {
90 + "cell_type": "code",
91 + "metadata": {},
92 + "execution_count": null,
93 + "outputs": [],
94 + "source": [
95 + "!python -m countrysense.clean --raw-dir \"{raw_dir}\""
96 + ]
97 + },
98 + {
99 + "cell_type": "markdown",
100 + "metadata": {},
101 + "source": [
102 + "## Treening\n",
103 + "\n",
104 + "Faas 1: backbone külmutatud, treenime ainult klassifitseerimispead (3 epohhi).\n",
105 + "Faas 2: viimased 4 transformeri plokki lahti, madal õppemäär (8 epohhi).\n",
106 + "\n",
107 + "Parim checkpoint salvestatakse faili `runs/clip_vit_b16/best.pt`."
108 + ]
109 + },
110 + {
111 + "cell_type": "code",
112 + "metadata": {},
113 + "execution_count": null,
114 + "outputs": [],
115 + "source": [
116 + "!python -m countrysense.train --config configs/default.yaml"
117 + ]
118 + },
119 + {
120 + "cell_type": "markdown",
121 + "metadata": {},
122 + "source": [
123 + "## Hindamine testihulgal"
124 + ]
125 + },
126 + {
127 + "cell_type": "code",
128 + "metadata": {},
129 + "execution_count": null,
130 + "outputs": [],
131 + "source": [
132 + "!python -m countrysense.evaluate --checkpoint runs/clip_vit_b16/best.pt"
133 + ]
134 + },
135 + {
136 + "cell_type": "code",
137 + "metadata": {},
138 + "execution_count": null,
139 + "outputs": [],
140 + "source": [
141 + "from IPython.display import Image as IPyImage\n",
142 + "IPyImage(\"runs/clip_vit_b16/confusion_matrix.png\")"
143 + ]
144 + },
145 + {
146 + "cell_type": "markdown",
147 + "metadata": {},
148 + "source": [
149 + "## Proovi ühe pildiga"
150 + ]
151 + },
152 + {
153 + "cell_type": "code",
154 + "metadata": {},
155 + "execution_count": null,
156 + "outputs": [],
157 + "source": [
158 + "import pandas as pd\n",
159 + "\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']}\""
163 + ]
164 + },
165 + {
166 + "cell_type": "markdown",
167 + "metadata": {},
168 + "source": [
169 + "## Checkpointi salvestamine Google Drive'i (valikuline)"
170 + ]
171 + },
172 + {
173 + "cell_type": "code",
174 + "metadata": {},
175 + "execution_count": null,
176 + "outputs": [],
177 + "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"
181 + ]
182 + }
183 + ]
184 +}
added requirements.txt +15 −0
@@ -0,0 +1,15 @@
1 +torch>=2.3
2 +torchvision>=0.18
3 +timm>=1.0.7
4 +pandas>=2.0
5 +numpy>=1.26
6 +scikit-learn>=1.4
7 +matplotlib>=3.8
8 +Pillow>=10.0
9 +ImageHash>=4.3
10 +tqdm>=4.66
11 +PyYAML>=6.0
12 +
13 +# valikuline:
14 +# gradio>=4.0 (app.py demo jaoks)
15 +# kagglehub>=0.3 (andmestiku allalaadimiseks Colabis)