Repository navigation
Expand file tree
/
Copy pathevaluate_ip.py
More file actions
701 lines (606 loc) · 32 KB
/
Copy pathevaluate_ip.py
File metadata and controls
701 lines (606 loc) · 32 KB
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
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
"""
evaluate_ip.py -- full monthly-schedule coverage evaluation.
Paper Sec. "Scalable Inference and Global Selection": a monthly timetable is
partitioned into temporal windows and each window is further divided into
connectivity-preserving chunks of at most n_max flight legs (subset_size /
config.EPISODE_MAX_FLIGHTS below), since a single rollout only covers a
bounded subset per episode. This script implements that full pipeline:
1. Split the full CSV into window_days-sized non-overlapping windows -> assign global flight IDs
2. Per window: partition into connectivity-preserving chunks -> stochastic
+ greedy policy rollouts per chunk -> accumulate the legal candidate
pool Cθ (keyed by global ID)
3. After all windows are processed, solve the restricted set-covering MIP
over Cθ (final selection) to cover the whole schedule
Notes:
- A pairing cannot be generated across a window boundary (a window-boundary limitation)
- IP scale: ~73,836 flights x pool pairings -> CBC can take hours (tune ip_time_limit)
"""
import sys
import os
import math
import random
import argparse
import torch
import pandas as pd
sys.path.insert(0, "RL")
DEVICE = torch.device("cpu")
from loader import load_flights_rolling, build_airport_map, bases_to_ids, sample_connected_subnet as sample_connected_subnet_std
from turkish.loader_turkish import (
parse_legs_dir, build_airport_map_turkish, load_flights_rolling_turkish,
sample_connected_subnet as sample_connected_subnet_turkish,
ZEREN_FEB_FILE, ZEREN_FEB_WINDOW,
)
from turkish.constraints_turkish import get_turkish_constraints as get_turkish_constraints_hb
from constraints import (
get_delta_constraints,
get_alaska_constraints,
get_jetblue_constraints,
FILM_CONSTRAINT_KEYS,
)
_GET_CONSTRAINT = {
"delta": get_delta_constraints,
"alaska": get_alaska_constraints,
"jetblue": get_jetblue_constraints,
"turkish": get_turkish_constraints_hb, # allows asymmetric HB1/HB2 termination
}
from model import FlightEncoder, PointerDecoder
from set_partition import solve_set_covering, solve_lp_relaxation
from utils import constraint_to_tensor, flights_to_tensors
from rollout import rollout_with_pairings, set_environment
from base_reach import build_base_reach
import config
# ── 0. Turkish window loader ─────────────────────────────────────────────────
def load_windows_turkish(turkish_df, airport_map, window_days=5):
"""Split Turkish .legs data into window_days-sized non-overlapping windows and assign global IDs."""
dates = sorted(turkish_df["dep_date_utc"].unique())
n_days = len(dates)
windows = []
global_offset = 0
for offset in range(0, n_days, window_days):
wf = load_flights_rolling_turkish(
window_days=window_days,
offset_days=offset,
airport_map=airport_map,
df=turkish_df,
)
for f in wf:
f["global_id"] = global_offset + f["id"]
global_offset += len(wf)
windows.append(wf)
print(
f" window offset={offset:2d}: {len(wf):5d}legs "
f"(global {global_offset - len(wf)} ~ {global_offset - 1})",
flush=True,
)
return windows, global_offset
# ── 1. Load the full dataset window by window, assigning global IDs ─────────
def load_windows_with_global_ids(data_path, airport_map, window_days=5, use_utc=False):
"""Split the full CSV into window_days-sized non-overlapping windows and assign global IDs.
use_utc: if True, anchor dep_time to UTC (see RL/loader.py) -- only enable
this when evaluating a checkpoint trained with the same option; using
it with an existing checkpoint puts the model out-of-distribution.
Returns:
windows : list of flight lists. Each flight gets a 'global_id' field.
n_total : total number of flights (= upper bound on global IDs)
"""
df = pd.read_csv(data_path)
df = df[["ORIGIN", "DEST", "CRS_DEP_TIME", "CRS_ARR_TIME", "CRS_ELAPSED_TIME", "FL_DATE"]].dropna()
df["FL_DATE"] = pd.to_datetime(df["FL_DATE"], format="mixed")
dates = sorted(df["FL_DATE"].unique())
n_days = len(dates)
windows = []
global_offset = 0
for offset in range(0, n_days, window_days):
wf = load_flights_rolling(
data_path,
window_days=window_days,
offset_days=offset,
airport_map=airport_map,
n_max=None,
df=df,
use_utc=use_utc,
)
for f in wf:
f["global_id"] = global_offset + f["id"]
global_offset += len(wf)
windows.append(wf)
print(
f" window offset={offset:2d}: {len(wf):5d}legs "
f"(global {global_offset - len(wf)} ~ {global_offset - 1})",
flush=True,
)
return windows, global_offset
# ── 2. Subset sampling within a window ───────────────────────────────────────
def sample_connected_subset(window_flights, subset_size, base_id, constraint):
"""Connectivity-aware subset sampling with random coverage guarantee.
Starts BFS from base-departing legs to preferentially select connectable
legs, but only fills BFS_RATIO of the subset via BFS; the rest is chosen
pure-random from the whole window.
Rationale: if BFS alone quickly fills the hub-and-spoke-dense region,
isolated flights with no connections are never included in any rollout,
capping pool coverage at ~85%. Guaranteeing 15% random inclusion gives
each flight an expected ~5 inclusions over 300 rollouts, keeping the
omission probability under 1%.
Args:
window_flights: list of all flights in the window
subset_size: number of legs to select (config.EPISODE_MAX_FLIGHTS)
base_id: crew base airport integer ID
constraint: constraint dict (min_conn, max_conn in hours)
"""
BFS_RATIO = 0.85 # 85% of the subset is BFS-connected (preserves RL multi-leg density)
min_conn = constraint.get("min_conn", 0.65) # hours
max_conn = constraint.get("max_conn", 9.0) # hours
# Index by origin airport
by_origin = {}
for f in window_flights:
by_origin.setdefault(f["origin"], []).append(f)
selected_ids = set()
selected = []
# BFS phase: fill only up to BFS_RATIO of the subset
bfs_quota = max(1, int(subset_size * BFS_RATIO))
base_departs = [f for f in window_flights if f["origin"] == base_id]
random.shuffle(base_departs)
queue = list(base_departs)
while queue and len(selected) < bfs_quota:
f = queue.pop(0)
if f["id"] in selected_ids:
continue
selected_ids.add(f["id"])
selected.append(f)
# Add legs connectable after this arrival to the queue
nexts = [
g for g in by_origin.get(f["dest"], [])
if g["id"] not in selected_ids
and min_conn <= g["dep_time"] - f["arr_time"] <= max_conn
]
random.shuffle(nexts)
queue.extend(nexts)
# Random phase: fill remaining slots pure-random from the whole window
# -> guarantees isolated flights get included with some probability every rollout
remaining = [f for f in window_flights if f["id"] not in selected_ids]
random.shuffle(remaining)
for f in remaining[:subset_size - len(selected)]:
selected_ids.add(f["id"])
selected.append(f)
selected = sorted(selected, key=lambda f: f["dep_time"])
for local_id, f in enumerate(selected):
f["local_id"] = local_id
return selected
# ── 3. Subset rollout -> global_id-keyed pairings ────────────────────────────
def rollout_subset_global(subset, constraint, encoder, decoder, max_time, greedy=False):
"""Roll out over a subset (n_max legs, with global_id) and return global_id-keyed pairings."""
local_flights = [{**f, "id": f["local_id"]} for f in subset]
origins, dests, dep_norm, arr_norm, fly_norm = flights_to_tensors(
local_flights, max_time, device=DEVICE
)
c_tensor = constraint_to_tensor(constraint, device=DEVICE)
with torch.no_grad():
encoded = encoder(origins, dests, dep_norm, arr_norm, fly_norm, c_tensor)
raw_pairings = rollout_with_pairings(
local_flights, constraint, encoder, decoder, encoded,
greedy=greedy, device=DEVICE,
)
id_map = {f["local_id"]: f["global_id"] for f in subset}
for p in raw_pairings:
p["legs"] = [id_map[leg] for leg in p["legs"]]
return raw_pairings
# ── 3-1. Partition a full window into connectivity-preserving chunks ────────
def partition_connected_chunks(window_flights, base_ids, chunk_size, connected_sampler):
"""Partition the whole window into connected-subnet chunks (Sec. "Scalable
Inference and Global Selection": divide each window into
connectivity-preserving chunks containing at most n_max flight legs).
Builds each chunk with the same sample_connected_subnet logic used during
training (RL/loader.py, RL/turkish/loader_turkish.py), repeating until
`remaining` is empty so every flight belongs to exactly one chunk
(100% coverage), keeping the same connectivity-density distribution seen at training time.
"""
remaining = list(window_flights)
chunks = []
while remaining:
for i, f in enumerate(remaining):
f["id"] = i
base_id = random.choice(base_ids)
chunk = connected_sampler(remaining, base_id, chunk_size)
if not chunk:
chunk = sorted(remaining, key=lambda f: f["dep_time"])[:chunk_size]
chosen_gids = {f["global_id"] for f in chunk}
remaining = [f for f in remaining if f["global_id"] not in chosen_gids]
chunks.append(sorted(chunk, key=lambda f: f["dep_time"]))
return chunks
# ── 4. Collect the pool across all windows ───────────────────────────────────
def collect_pool_full(windows, base_ids, constraint, encoder, decoder,
n_rollouts_per_chunk=5,
subset_size=config.EPISODE_MAX_FLIGHTS,
connected_sampler=sample_connected_subnet_std,
airline="delta",
require_base_return=False):
"""Roll out over all windows to build the global-ID-keyed candidate pool Cθ.
Paper Sec. "Scalable Inference and Global Selection": "For each chunk, we
perform multiple stochastic rollouts and one greedy rollout... Candidates
from all chunks are merged into a global pool Cθ." Each window is split
into connectivity-preserving chunks via connected_sampler (the same
sample_connected_subnet used during training); each chunk gets
n_rollouts_per_chunk stochastic rollouts plus 1 greedy rollout.
Partitioning repeats until `remaining` is empty, so every flight is
included in at least one rollout (guaranteeing 100% coverage
opportunity) while preserving the same connectivity density seen during training.
airline="turkish" allows the two bases HB1/HB2 to substitute for each
other (environment_turkish.py) -- rollout.py's p["ends_at_base"] only
checks same-base return against the single base assigned to that
rollout, so it would incorrectly reject a valid HB1->HB2 cross-return.
For turkish only, validity is instead determined by checking whether the
actual first/last leg's origin/dest lie in base_id_set (all of HB1, HB2).
"""
pool = {}
covered_global = set()
max_time = 5 * 24.0
base_id_set = set(base_ids)
for w_idx, window_flights in enumerate(windows):
if not window_flights:
continue
window_all_ids = set(f["global_id"] for f in window_flights)
window_covered = set()
chunks = partition_connected_chunks(window_flights, base_ids, subset_size, connected_sampler)
# Guarantee a base-departing leg in each chunk: if missing, inject a
# copy of the nearest base-departing leg from the window
# (connected-subnet partitioning usually includes one, but tail
# chunks etc. can be an exception).
for c_idx, chunk in enumerate(chunks):
if not any(f["origin"] in base_id_set for f in chunk):
chunk_gids = {f["global_id"] for f in chunk}
candidates = [f for f in window_flights
if f["origin"] in base_id_set and f["global_id"] not in chunk_gids]
if candidates:
inject = min(candidates,
key=lambda f: abs(f["dep_time"] - chunk[0]["dep_time"]))
new_chunk = sorted([{**inject}] + list(chunk[:-1]),
key=lambda f: f["dep_time"])
chunks[c_idx] = new_chunk
print(f"\n[Window {w_idx + 1}/{len(windows)}] {len(window_flights)} legs -> {len(chunks)} chunks", flush=True)
rollout_count = 0
for c_idx, chunk in enumerate(chunks):
for local_id, f in enumerate(chunk):
f["local_id"] = local_id
chunk_by_gid = {f["global_id"]: f for f in chunk}
def _pairing_valid(p, _chunk_by_gid=chunk_by_gid):
if airline != "turkish":
return p["ends_at_base"]
# turkish: HB1->HB2 and HB2->HB1 are also valid -- rollout.py's
# ends_at_base (single-episode_base same-base check) would
# reject this cross-return, so re-derive validity by
# comparing the actual first/last leg origin/dest against the
# full base_id_set.
first = _chunk_by_gid.get(p["legs"][0])
last = _chunk_by_gid.get(p["legs"][-1])
return (first is not None and last is not None
and first["origin"] in base_id_set and last["dest"] in base_id_set)
base_id = random.choice(base_ids)
c_b = {**constraint, "base_airport": base_id}
if require_base_return:
# rollout.py supports base rotation (switch to another base
# once the current base's departing legs are exhausted) and
# salvage (on a dead end, keep only the prefix that ends at
# base and return the rest); pass base_ids/strict_base_start
# to activate that path.
c_b["base_ids"] = base_ids
c_b["strict_base_start"] = True
c_b["require_base_return"] = True
# rollout_subset_global remaps flight["id"] to local_id before
# rolling out (mask/step indexing breaks under global IDs), so
# reachability must also be computed against local_id to match
# the IDs actually looked up during rollout.
_local_flights = [{**f, "id": f["local_id"]} for f in chunk]
c_b["_base_reach"] = build_base_reach(_local_flights, base_id, c_b)
for _ in range(n_rollouts_per_chunk):
try:
pairings = rollout_subset_global(chunk, c_b, encoder, decoder, max_time, greedy=False)
except Exception as e:
print(f" [warn] stochastic rollout failed (chunk={c_idx}): {e}", flush=True)
continue
for p in pairings:
# Exclude pairings that don't return to base from both the
# pool and coverage counts -- including them in coverage
# would create "phantom" coverage that the IP can never
# actually select, so window_covered/covered_global must
# be filtered the same way.
if not _pairing_valid(p):
continue
key = tuple(sorted(p["legs"]))
if key not in pool or p["cost"] < pool[key]["cost"]:
pool[key] = p
window_covered.update(p["legs"])
covered_global.update(p["legs"])
rollout_count += 1
try:
pairings = rollout_subset_global(chunk, c_b, encoder, decoder, max_time, greedy=True)
except Exception as e:
print(f" [warn] greedy rollout failed (chunk={c_idx}): {e}", flush=True)
continue
for p in pairings:
if not _pairing_valid(p):
continue
key = tuple(sorted(p["legs"]))
if key not in pool or p["cost"] < pool[key]["cost"]:
pool[key] = p
window_covered.update(p["legs"])
covered_global.update(p["legs"])
rollout_count += 1
print(
f" chunk {c_idx + 1}/{len(chunks)} done "
f"(cumulative rollouts={rollout_count}, pool={len(pool)}, "
f"window covered={len(window_covered)}/{len(window_all_ids)})",
flush=True,
)
print(
f" {rollout_count} total rollouts: "
f"window covered {len(window_covered)}/{len(window_all_ids)} legs, "
f"pool={len(pool)}",
flush=True,
)
uncov = len(window_all_ids - window_covered)
if uncov > 0:
print(f" uncovered: {uncov} legs (reported as uncoverable by the IP)", flush=True)
total_flights = sum(len(w) for w in windows)
print(f"\ntotal pool: {len(pool)} pairings")
print(f"total coverage: {len(covered_global)}/{total_flights} legs")
return list(pool.values()), covered_global
# ── 5. Main evaluation function ──────────────────────────────────────────────
def evaluate_full(
checkpoint_path,
airline="delta",
data_path=None,
n_rollouts_per_chunk=5,
window_days=5,
subset_size=config.EPISODE_MAX_FLIGHTS,
bases=None,
ip_time_limit=3600,
lambda_dh=1.0,
device="cpu",
turkish_files=None,
use_utc=False,
use_wandb=False,
wandb_project="ASCP-2026-paper",
compute_gap=False,
seed=None,
require_base_return=False,
):
"""Full flight-coverage evaluation. Uses config.AIRLINE_DATA[airline] if data_path is unset.
For small-scale data (e.g. a one-week sample):
data_path=<sample.csv>, window_days=1, n_rollouts_per_chunk=3
If use_wandb=True, logs the eval config + console output + final result
metrics to wandb (job_type="eval" -- kept as a separate run from training curves).
If seed is given, fixes the random/torch RNG so different checkpoints are
evaluated on the same window partitioning and the same rollout sampling,
enabling paired comparison (differences between checkpoints then come
only from policy differences, not evaluation randomness).
"""
if seed is not None:
random.seed(seed)
torch.manual_seed(seed)
global DEVICE
DEVICE = torch.device(device)
if device == "cpu":
# By default torch claims as many threads per process as there are
# physical cores, which causes heavy CPU contention when evaluating
# multiple checkpoints in parallel (and with concurrent GPU training
# processes). Some torch builds ignore env-level limits like
# OMP_NUM_THREADS, so pin it explicitly here as well.
torch.set_num_threads(int(os.environ.get("OMP_NUM_THREADS", 4)))
set_environment(airline)
wandb_run = None
if use_wandb:
import wandb
wandb_run = wandb.init(
project=wandb_project,
job_type="eval",
name=f"eval-{airline}-{os.path.basename(checkpoint_path)}",
config=dict(
checkpoint=checkpoint_path, airline=airline,
subset_size=subset_size, window_days=window_days,
n_rollouts_per_chunk=n_rollouts_per_chunk,
ip_time_limit=ip_time_limit, lambda_dh=lambda_dh,
use_utc=use_utc,
),
)
if data_path is None:
data_path = config.AIRLINE_DATA[airline]
if bases is None:
bases = config.AIRLINE_BASES[airline]
# Load the checkpoint first to check its vocab size -- a multi-airline
# model (n_airports=168) needs the merged airport map; building it from a
# single-airline map would cause an embedding-index mismatch.
ckpt = torch.load(checkpoint_path, map_location=DEVICE, weights_only=True)
n_airports = ckpt.get("n_airports",
ckpt["encoder"]["airport_emb.weight"].shape[0])
_turkish_df = None
if airline == "turkish":
# If turkish_files is unset, default to the Zeren Feb benchmark
# window (15,742 legs, 0.03% off the target 15,738)
if turkish_files is None:
_turkish_df = parse_legs_dir(data_path, files=[ZEREN_FEB_FILE], date_range=ZEREN_FEB_WINDOW)
else:
_turkish_df = parse_legs_dir(data_path, files=turkish_files)
airport_map = build_airport_map_turkish(df=_turkish_df)
else:
if n_airports > 145:
# Turkish (.legs directory) can't be processed by the BTS CSV loader -> exclude
map_paths = [v for k, v in config.AIRLINE_DATA.items() if k != "turkish"]
else:
map_paths = data_path
airport_map = build_airport_map(map_paths)
base_ids = bases_to_ids(list(bases), airport_map)
encoder = FlightEncoder(n_airports=n_airports, constraint_dim=len(FILM_CONSTRAINT_KEYS)).to(DEVICE)
# Auto-detect the checkpoint's state_vec dimension (older checkpoints used
# fewer scalars than the current state_to_vec)
airport_emb_dim = encoder.airport_emb.embedding_dim
ckpt_state_dim = ckpt["decoder"]["state_mlp.0.weight"].shape[1]
n_scalars = ckpt_state_dim - airport_emb_dim * 2 - len(FILM_CONSTRAINT_KEYS)
decoder = PointerDecoder(constraint_dim=len(FILM_CONSTRAINT_KEYS), airport_emb_dim=airport_emb_dim, n_scalars=n_scalars).to(DEVICE)
encoder.load_state_dict(ckpt["encoder"])
decoder.load_state_dict(ckpt["decoder"])
encoder.eval()
decoder.eval()
if airline == "turkish":
constraint = _GET_CONSTRAINT[airline](base_ids[0], base_ids=base_ids)
else:
constraint = _GET_CONSTRAINT[airline](base_ids[0])
print(f"\nLoading full dataset ({airline}, window_days={window_days})...", flush=True)
if airline == "turkish":
windows, n_total = load_windows_turkish(_turkish_df, airport_map, window_days)
else:
windows, n_total = load_windows_with_global_ids(data_path, airport_map, window_days, use_utc=use_utc)
print(f"total {n_total} legs, {len(windows)} windows", flush=True)
connected_sampler = sample_connected_subnet_turkish if airline == "turkish" else sample_connected_subnet_std
_hard_mask = require_base_return
if _hard_mask:
print("\n[base-return] decode-time hard mask ON (includes reachability pruning)", flush=True)
if airline == "turkish":
print(" [note] For turkish, HB1<->HB2 cross-return is not enforced; the hard "
"mask only enforces single-base return to whichever base the pairing "
"actually departed from (a stricter subset).",
flush=True)
print(f"\nCollecting pool (rollouts/chunk={n_rollouts_per_chunk}, subset={subset_size})...", flush=True)
with torch.no_grad():
pool, covered = collect_pool_full(
windows, base_ids, constraint, encoder, decoder,
n_rollouts_per_chunk=n_rollouts_per_chunk,
subset_size=subset_size,
connected_sampler=connected_sampler,
airline=airline,
require_base_return=_hard_mask,
)
print(f"\nSolving IP (n_flights={n_total}, pool={len(pool)}, time_limit={ip_time_limit}s, lambda_dh={lambda_dh})...", flush=True)
result = solve_set_covering(pool, n_flights=n_total, time_limit=ip_time_limit, lambda_dh=lambda_dh, verbose=True)
print("IP solve complete", flush=True)
gap_pct = None
if compute_gap:
print(f"\nSolving LP relaxation (for Gap%, pool={len(pool)})...", flush=True)
lp_result = solve_lp_relaxation(pool, lambda_dh=lambda_dh)
if lp_result is not None and lp_result["lp_value"]:
gap_pct = (result["mip_obj"] - lp_result["lp_value"]) / lp_result["lp_value"] * 100
else:
print(" [warn] LP relaxation failed to solve -- cannot compute Gap%")
sel = result["selected"]
fly_total = sum(p["fly"] for p in sel) if sel else 0.0
raw_dead_total = sum(p.get("dead_time", p["cost"]) for p in sel) if sel else 0.0
legs_total = sum(p.get("n_legs", len(p["legs"])) for p in sel) if sel else 0
duties_total = sum(p.get("n_duties", 1) for p in sel) if sel else 0
man_days = sum(math.ceil(p["elapsed"] / 24.0) for p in sel) if sel else 0
avg_legs = legs_total / len(sel) if sel else 0.0
avg_duties = duties_total / len(sel) if sel else 0.0
# FTC reflects only within-duty gaps (excludes overnight excess); cost is
# left as-is (preserves the ManDays incentive)
intra_gap_total = sum(p.get("intra_duty_gap", 0.0) for p in sel) if sel else 0.0
inter_excess_total = sum(p.get("inter_duty_excess", 0.0) for p in sel) if sel else 0.0
# Total dead time is reported on the same basis as FTC (within-duty gaps
# only, excludes overnight excess) -- raw_dead_total (the cost-computation
# basis, includes overnight excess) is shown separately for reference
dead_total = intra_gap_total
ftc = intra_gap_total / fly_total * 100 if fly_total > 0 else 0.0
print()
print("=" * 60)
print(f"Results (covering all {n_total} legs)")
print("=" * 60)
print(f" n pairings: {result['n_pairings']}")
print(f" ManDays: {man_days}")
print(f" coverage: {result['coverage'] * 100:.1f}%")
print(f" uncoverable: {result['uncoverable']} legs")
print(f" deadhead: {result['deadhead_count']} legs")
print(f" fly time: {fly_total:.2f}h")
print(f" dead time (within-duty gaps only, excl. overnight): {dead_total:.2f}h")
print(f" (ref) raw dead time (incl. overnight excess, cost-computation basis): {raw_dead_total:.2f}h")
print(f" - within-duty connection gap: {intra_gap_total:.2f}h ({intra_gap_total/raw_dead_total*100 if raw_dead_total>0 else 0:.1f}%)")
print(f" - inter-duty excess wait (>min_rest): {inter_excess_total:.2f}h ({inter_excess_total/raw_dead_total*100 if raw_dead_total>0 else 0:.1f}%)")
print(f" FTC: {ftc:.2f}%")
print(f" avg legs/pairing: {avg_legs:.2f}")
print(f" avg duties/pairing:{avg_duties:.2f}")
print(f" IP status: {result['status']}")
if gap_pct is not None:
print(f" Gap% (MIP vs LP): {gap_pct:.3f}%")
if wandb_run is not None:
import wandb
wandb.log({
"n_pairings": result["n_pairings"],
"man_days": man_days,
"coverage": result["coverage"] * 100,
"uncoverable": result["uncoverable"],
"deadhead": result["deadhead_count"],
"fly_time": fly_total,
"dead_time": dead_total,
"raw_dead_time": raw_dead_total,
"intra_duty_gap": intra_gap_total,
"inter_duty_excess": inter_excess_total,
"ftc": ftc,
"avg_legs": avg_legs,
"avg_duties": avg_duties,
"ip_status": result["status"],
"gap_pct": gap_pct,
})
wandb.finish()
result["gap_pct"] = gap_pct
return result
# ── CLI ────────────────────────────────────────────────────────────────────────
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Full monthly-schedule flight-coverage evaluation")
parser.add_argument("checkpoint", help="Checkpoint file path (e.g. checkpoints/jbkwcdk3/phase2_best.pt)")
parser.add_argument("--airline", default="delta", choices=["delta", "alaska", "jetblue", "turkish"])
parser.add_argument("--data-path", default=None,
help="CSV path. Uses config.AIRLINE_DATA[airline] if unset. "
"Set this for small-scale sample evaluation (e.g. RL/data/sample_DL_*.csv)")
parser.add_argument("--n-rollouts-per-chunk", type=int, default=5,
help="Stochastic rollouts per chunk. Each window is split into sequential subset_size-sized chunks (default: 5)")
parser.add_argument("--window-days", type=int, default=5,
help="Window size in days. 1 is recommended for small-scale (1-week) data (default: 5)")
parser.add_argument("--subset-size", type=int, default=config.EPISODE_MAX_FLIGHTS,
help=f"Flights per rollout (default: {config.EPISODE_MAX_FLIGHTS})")
parser.add_argument("--ip-time-limit", type=int, default=3600,
help="CBC solver time limit in seconds (default: 3600)")
parser.add_argument("--lambda-dh", type=float, default=1.0,
help="DH penalty weight (default: 1.0)")
parser.add_argument("--device", default="cpu")
parser.add_argument("--turkish-files", nargs="+", default=None,
help="Turkish only. List of .legs file names to use. Defaults to the "
"Zeren Feb benchmark window (tt201402.legs, 2/1-3/8, 15,742 legs) "
"if unset. If given explicitly, uses those files in full with no date filter.")
parser.add_argument("--use-utc", action="store_true",
help="Anchor dep_time as absolute UTC time. Only evaluate with this flag "
"for checkpoints trained with --use-utc -- enabling it for an existing "
"checkpoint puts the model out-of-distribution")
parser.add_argument("--wandb", action="store_true",
help="Log the eval config + console output + final result metrics to wandb (job_type=eval)")
parser.add_argument("--wandb-project", default="ASCP-2026-paper")
parser.add_argument("--compute-gap", action="store_true",
help="After solving the MIP, also solve the LP relaxation over the same "
"pool to compute Gap%%=(MIP_obj-LP_obj)/LP_obj*100 (same definition as "
"Tahir et al. Table 6). Off by default since the LP adds extra time on large pools")
parser.add_argument("--seed", type=int, default=None,
help="Fix the random/torch RNG -- set this to run a paired comparison of "
"multiple checkpoints against the same evaluation instance (e.g. the "
"same seed for every ON/OFF checkpoint)")
parser.add_argument("--require-base-return", action="store_true",
help="Enable the decode-time hard mask -- masks any leg that would make "
"base return infeasible during rollout, and forbids EndPairing away from the base.")
args = parser.parse_args()
ckpt = args.checkpoint
if not os.path.exists(ckpt):
candidate = os.path.join("checkpoints", ckpt)
if os.path.exists(candidate):
ckpt = candidate
evaluate_full(
checkpoint_path=ckpt,
airline=args.airline,
data_path=args.data_path,
n_rollouts_per_chunk=args.n_rollouts_per_chunk,
window_days=args.window_days,
subset_size=args.subset_size,
ip_time_limit=args.ip_time_limit,
lambda_dh=args.lambda_dh,
device=args.device,
turkish_files=args.turkish_files,
use_utc=args.use_utc,
use_wandb=args.wandb,
wandb_project=args.wandb_project,
compute_gap=args.compute_gap,
seed=args.seed,
require_base_return=args.require_base_return,
)