File size: 3,002 Bytes
143710c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7b4e05b
 
143710c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
"""Config loading utility.

Every experiment (architecture choice, hyperparameters, paths) is described
by a YAML file under configs/. This module is the single place that reads
those files into a plain dict, so every script (split, vocab, features,
train, evaluate, predict) loads config the same way.
"""
from __future__ import annotations

import copy
from pathlib import Path
from typing import Any

import yaml

DEFAULT_CONFIG: dict[str, Any] = {
    "seed": 42,
    "paths": {
        "raw_images_dir": "data/raw/images",
        "captions_file": "data/raw/captions.txt",
        "processed_dir": "data/processed",
        "features_dir": "data/features",
        "models_dir": "models",
    },
    "split": {
        "train_ratio": 0.8,
        "val_ratio": 0.1,
        "test_ratio": 0.1,
    },
    "vocab": {
        "min_freq": 5,
        "max_len": 35,
    },
    "encoder": {
        "type": "resnet50",
        "feature_dim": 2048,
        "freeze": True,
    },
    "decoder": {
        "type": "lstm",
        "embed_dim": 256,
        "hidden_dim": 512,
        "num_layers": 1,
        "dropout": 0.0,
    },
    "training": {
        "batch_size": 32,
        "epochs": 30,
        "lr": 1e-3,
        "optimizer": "adam",
        "weight_decay": 0.0,
        "grad_clip_norm": 0.0,
        "lr_scheduler_patience": 2,
        "lr_scheduler_factor": 0.5,
        "early_stop_patience": 5,
        "num_workers": 2,
        "device": "auto",  # "auto" | "cpu" | "cuda"
    },
    "inference": {
        "decoding": "greedy",  # "greedy" | "beam"
        "beam_width": 3,
    },
    "run_name": "base_resnet_lstm",
}


def _deep_update(base: dict, override: dict) -> dict:
    """Recursively merge override into base (override wins), without mutating inputs."""
    result = copy.deepcopy(base)
    for key, value in override.items():
        if isinstance(value, dict) and isinstance(result.get(key), dict):
            result[key] = _deep_update(result[key], value)
        else:
            result[key] = value
    return result


def load_config(path: str | Path) -> dict[str, Any]:
    """Load a YAML config and merge it on top of DEFAULT_CONFIG.

    Any key not specified in the YAML file falls back to DEFAULT_CONFIG,
    so config files only need to declare what they change from baseline.
    """
    path = Path(path)
    if not path.exists():
        raise FileNotFoundError(f"Config file not found: {path}")

    with open(path, "r") as f:
        user_config = yaml.safe_load(f) or {}

    config = _deep_update(DEFAULT_CONFIG, user_config)
    config["_config_path"] = str(path)
    return config


def save_config(config: dict[str, Any], path: str | Path) -> None:
    """Persist a resolved config dict to YAML (e.g. alongside a checkpoint)."""
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    clean = {k: v for k, v in config.items() if not k.startswith("_")}
    with open(path, "w") as f:
        yaml.safe_dump(clean, f, sort_keys=False)