File size: 3,929 Bytes
4198d45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
#!/usr/bin/env python3
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""Run released pre-decoders on the fixed training-axis OOD grid."""

from __future__ import annotations

import argparse
import sys
from pathlib import Path

CODE_ROOT = Path(__file__).resolve().parents[1]
if str(CODE_ROOT) not in sys.path:
    sys.path.insert(0, str(CODE_ROOT))

from scripts.experiments.unknown_noise.generate_unknown_axismix_grid_u1p2_5p0_configs import (  # noqa: E402
    write_axismix_grid_configs,
)
from scripts.qadapt_example_utils import (  # noqa: E402
    InferenceJob,
    add_common_inference_args,
    build_paired_command,
    parse_gpus,
    run_jobs,
)


PAPER_DISTANCES = (7, 9)
PAPER_MULTIPLIERS = (1.2, 1.5, 2.0, 2.5, 3.0)


def parse_distances(value: str) -> list[int]:
    result = [int(item.strip()) for item in value.split(",") if item.strip()]
    if not result or result != sorted(set(result)):
        raise argparse.ArgumentTypeError(
            "distances must be a non-empty, increasing comma-separated list"
        )
    return result


def parse_multipliers(value: str) -> list[float]:
    result = [float(item.strip()) for item in value.split(",") if item.strip()]
    if not result or result != sorted(set(result)) or any(item <= 0 for item in result):
        raise argparse.ArgumentTypeError(
            "multipliers must be a non-empty, increasing comma-separated list "
            "of positive numbers"
        )
    return result


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--distances",
        type=parse_distances,
        default=list(PAPER_DISTANCES),
        help="Comma-separated distances; defaults to the paper's d=7,9 grid.",
    )
    parser.add_argument("--n-rounds", type=int, default=9)
    parser.add_argument(
        "--multipliers",
        type=parse_multipliers,
        default=list(PAPER_MULTIPLIERS),
        help="Comma-separated OOD multipliers; defaults to the paper's 1.2--3.0 grid.",
    )
    parser.add_argument(
        "--generated-config-dir",
        type=Path,
        default=Path("outputs/generated_configs/ood"),
    )
    parser.add_argument(
        "--manifest",
        type=Path,
        default=Path("outputs/generated_configs/ood/manifest.json"),
    )
    add_common_inference_args(
        parser,
        default_output_dir=Path("outputs/examples/released_models/ood"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    _, manifest = write_axismix_grid_configs(
        base_config="conf/examples/qadapt/config_qadapt_t0_base.yaml",
        output_dir=args.generated_config_dir,
        manifest=args.manifest,
        grid_multipliers=args.multipliers,
    )
    jobs = []
    for distance in args.distances:
        for environment in manifest["environments"]:
            config_file = args.generated_config_dir / environment["config_filename"]
            label = (
                f"d{distance}_{environment['env_key']}_"
                f"{environment['multiplier_key']}"
            )
            output_path = args.output_dir / f"d{distance}" / f"{label}.json"
            jobs.append(
                InferenceJob(
                    label=label,
                    command=build_paired_command(
                        args,
                        config_file=config_file,
                        output_path=output_path,
                        distance=distance,
                        n_rounds=args.n_rounds,
                    ),
                    output_path=output_path,
                )
            )
    run_jobs(
        jobs,
        gpus=parse_gpus(args.gpus),
        parallelism=args.parallelism,
        resume=args.resume,
        dry_run=args.dry_run,
    )


if __name__ == "__main__":
    main()