Compare commits

..

No commits in common. "292f1ea6067108e0df35936f23c431e2b7966765" and "1da467d1d6438385ce451bd75e53c58d628c46c6" have entirely different histories.

56 changed files with 323 additions and 8706 deletions

View file

@ -1,10 +1,9 @@
/artifacts /artifacts
/data/raw /data/raw
/data/processed
/.venv /.venv
/notebooks /notebooks
__pycache__ .env
*.pyc
HF_TOKEN HF_TOKEN
WANDB_API_KEY WANDB_API_KEY
.env __pycache__
*.pyc

View file

@ -1,5 +1,5 @@
[run] [run]
name = "model_class_frontier_7gb_01_film_fourier_inr" name = "aggressive_smoke"
seed = 20260723 seed = 20260723
artifact_dir = "artifacts/current_run/training_runs" artifact_dir = "artifacts/current_run/training_runs"
@ -10,14 +10,9 @@ val_cases = 3
test_cases = 2 test_cases = 2
points_per_case = 999999999 points_per_case = 999999999
batch_size = 4096 batch_size = 4096
source = "huggingface"
hf_repo_id = "zacheryasc/airfrans-processed"
hf_repo_type = "dataset"
hf_path_prefix = "processed/full"
cache_dir = "artifacts/data_cache/airfrans_processed"
[model] [model]
type = "film_fourier_inr" type = "film_fourier_mlp"
hidden_width = 4096 hidden_width = 4096
depth = 12 depth = 12
activation = "gelu" activation = "gelu"
@ -54,12 +49,4 @@ max_grad_norm = 1.0
backend = "wandb" backend = "wandb"
entity = "zacheryasc-personal" entity = "zacheryasc-personal"
project = "airfRANS-model-sweep" project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01" tags = ["airfrans", "remote", "aggressive-smoke"]
tags = ["airfrans", "model-class-frontier", "7gb-subset", "flop-par", "hf-checkpoints"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01"
private = false

View file

@ -1,73 +0,0 @@
[run]
name = "full_airfrans_incumbent_70gb_01"
seed = 20260723
artifact_dir = "artifacts/current_run/training_runs"
[data]
root = "artifacts/data_cache/airfrans_streaming_processed/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
batch_size = 4096
source = "public_zip_streaming"
public_source_url = "https://data.isir.upmc.fr/extrality/NeurIPS_2022/OF_dataset.zip"
hf_repo_id = "zacheryasc/airfrans-processed"
hf_repo_type = "dataset"
hf_path_prefix = "processed/full"
cache_dir = "artifacts/data_cache/airfrans_streaming_processed/processed/full"
streaming_scratch_dir = "artifacts/data_cache/airfrans_streaming_processed/raw_scratch"
streaming_cache_max_bytes = 68719476736
streaming_cache_high_water_bytes = 51539607552
streaming_cache_low_water_bytes = 34359738368
streaming_queue_max_cases = 2
streaming_upload_processed = true
streaming_upload_batch_size = 16
[model]
type = "film_fourier_inr"
hidden_width = 4096
depth = 12
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
fourier_scales = [1.0, 2.0, 4.0, 8.0, 16.0, 32.0]
condition_width = 1024
condition_depth = 3
condition_dim = 512
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_airfrans_incumbent_70gb_01"
tags = ["airfrans", "full-data-frontier", "70gb", "incumbent", "hf-checkpoints"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_airfrans_incumbent_70gb_01"
private = false

View file

@ -1,61 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_deeponet_branch_trunk"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/deeponet_branch_trunk"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "deeponet_branch_trunk"
hidden_width = 1024
depth = 8
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
fourier_scales = [1.0, 2.0, 4.0, 8.0, 16.0, 32.0]
condition_width = 1024
condition_depth = 3
condition_dim = 512
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "deeponet_branch_trunk"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/deeponet_branch_trunk"
private = false

View file

@ -1,61 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_film_fourier_inr"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/film_fourier_inr"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "film_fourier_inr"
hidden_width = 4096
depth = 12
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
fourier_scales = [1.0, 2.0, 4.0, 8.0, 16.0, 32.0]
condition_width = 1024
condition_depth = 3
condition_dim = 512
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "film_fourier_inr"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/film_fourier_inr"
private = false

View file

@ -1,61 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_meshgraphnet_or_point_transformer_local"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/meshgraphnet_or_point_transformer_local"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "meshgraphnet_or_point_transformer_local"
hidden_width = 512
depth = 4
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
condition_width = 512
condition_depth = 3
condition_dim = 512
neighbors = 8
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "meshgraphnet_or_point_transformer_local"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/meshgraphnet_or_point_transformer_local"
private = false

View file

@ -1,62 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_nerf_cfd_multires"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/nerf_cfd_multires"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "nerf_cfd_multires"
hidden_width = 1024
depth = 8
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
condition_width = 512
condition_depth = 3
condition_dim = 512
encoding_levels = 16
features_per_level = 2
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "nerf_cfd_multires"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/nerf_cfd_multires"
private = false

View file

@ -1,63 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_point_context_perceiver"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/point_context_perceiver"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "point_context_perceiver"
hidden_width = 512
depth = 4
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
condition_width = 512
condition_depth = 3
condition_dim = 512
context_points = 512
latent_width = 512
attention_depth = 4
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "point_context_perceiver"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/point_context_perceiver"
private = false

View file

@ -1,61 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_raster_fno_unet"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/raster_fno_unet"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "raster_fno_unet"
hidden_width = 512
depth = 6
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
condition_width = 512
condition_depth = 3
condition_dim = 512
grid_resolution = 128
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "raster_fno_unet"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/raster_fno_unet"
private = false

View file

@ -1,61 +0,0 @@
[run]
name = "model_class_frontier_7gb_01_siren_conditioned_inr"
seed = 20260725
artifact_dir = "artifacts/current_run/training_runs/siren_conditioned_inr"
[data]
root = "data/processed/full"
train_cases = 45
val_cases = 3
test_cases = 2
points_per_case = 999999999
batch_size = 4096
source = "local"
[model]
type = "siren_conditioned_inr"
hidden_width = 1024
depth = 6
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
condition_width = 512
condition_depth = 3
condition_dim = 512
siren_omega0 = 30.0
[optim]
lr = 0.0001
weight_decay = 0.0001
steps = 5000
log_interval = 500
[device]
type = "cuda"
allow_cpu_fallback = false
benchmark_kernels = true
[loss]
type = "normalized_mse"
[precision]
dtype = "bf16"
[checkpoint]
interval_seconds = 1800
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "model_class_frontier_7gb_01"
tags = ["airfrans", "model-class-frontier", "7gb-subset", "tarball-data", "siren_conditioned_inr"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "model_class_frontier_7gb_01/siren_conditioned_inr"
private = false

View file

@ -1,90 +0,0 @@
[run]
name = "full_airfrans_incumbent_70gb_01"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
[provider]
kind = "vastai"
disk_gb = 192
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/full_airfrans_incumbent_70gb.toml'); assert c.data.source == 'public_zip_streaming'; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.streaming_cache_high_water_bytes < c.data.streaming_cache_max_bytes; assert c.data.streaming_cache_low_water_bytes < c.data.streaming_cache_high_water_bytes; print('data_source=' + c.data.source + ' public_url=' + str(c.data.public_source_url) + ' split=' + str((c.data.train_cases, c.data.val_cases, c.data.test_cases)) + ' cache_high_water=' + str(c.data.streaming_cache_high_water_bytes))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/full_airfrans_incumbent_70gb.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
"streaming_events.jsonl",
"streaming_state.json",
"streaming_summary.json",
"processed_upload_manifest.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_deeponet_branch_trunk"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/deeponet_branch_trunk.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/deeponet_branch_trunk.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_film_fourier_inr"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/film_fourier_inr.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/film_fourier_inr.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_meshgraphnet_or_point_transformer_local"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/meshgraphnet_or_point_transformer_local.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/meshgraphnet_or_point_transformer_local.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_nerf_cfd_multires"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/nerf_cfd_multires.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/nerf_cfd_multires.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_point_context_perceiver"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/point_context_perceiver.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/point_context_perceiver.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_raster_fno_unet"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/raster_fno_unet.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/raster_fno_unet.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,89 +0,0 @@
[run]
name = "model_zoo_7gb_01_siren_conditioned_inr"
timeout_minutes = 720
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2
[provider]
kind = "vastai"
disk_gb = 128
max_price_per_hour = 0.80
image = "vastai/base:0.0.2"
[provider.gpu]
name = "RTX 4090"
count = 1
min_vram_gb = 20
[selection]
min_reliability = 0.95
min_down_mbps = 100
min_up_mbps = 25
require_verified = true
blocked_geos = ["CN"]
blacklist_hosts = [59017, 1647, 92578, 1276, 75481, 1256, 85323, 34031]
drop_cheap_frac = 0.30
image_size_gb = 5.0
base_url = "https://cloud.vast.ai"
[workspace]
workdir = "."
exclude = [
"/artifacts",
"/data/raw",
"/data/processed",
"/.venv",
"/notebooks",
"__pycache__",
"*.pyc",
]
[bootstrap]
command = """
uv sync --no-dev
uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=torch.cuda.device_count(); print('torch_cuda_available=' + str(ok)); print('torch_cuda_version=' + str(torch.version.cuda)); print('torch_device_count=' + str(count)); print('torch_device_name=' + (torch.cuda.get_device_name(0) if ok and count else 'none')); assert ok, 'torch CUDA unavailable'"
"""
[data]
validation_command = """
mkdir -p data/processed
tar -xzf data/airfrans_processed_full_50cases.tar.gz -C data/processed
uv run --no-dev python -c "from pathlib import Path; files=list(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print('tarball_processed_full_cases=' + str(len(files)))"
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/model_zoo_7gb/siren_conditioned_inr.toml'); assert c.data.source == 'local'; assert c.data.train_cases == 45; assert c.data.val_cases == 3; assert c.data.test_cases == 2; print('model_family=' + c.model.type + ' data_source=' + c.data.source)"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_7gb/siren_conditioned_inr.toml --artifact-dir artifacts/current_run --run-id "$AIRFRANS_REMOTE_RUN_ID"
"""
artifact_dir = "artifacts/current_run"
heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts]
mode = "object_store_upload"
required = [
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
]
[cleanup]
on_success = "sky_down"
on_failure = "sky_down"

View file

@ -1,5 +1,5 @@
[run] [run]
name = "model_class_frontier_7gb_01_film_fourier_inr" name = "airfrans-aggressive-smoke"
timeout_minutes = 360 timeout_minutes = 360
local_artifact_dir = "artifacts/remote_runs" local_artifact_dir = "artifacts/remote_runs"
max_attempts = 2 max_attempts = 2
@ -45,7 +45,7 @@ uv run --no-dev python -c "import torch; assert torch.cuda.is_available(); print
[data] [data]
validation_command = """ validation_command = """
uv run --no-dev python -c "from airfrans_frontier.training.config import load_training_config; c=load_training_config('configs/aggressive_smoke.toml'); assert c.data.source == 'huggingface'; print('data_source=' + c.data.source + ' repo=' + str(c.data.hf_repo_id))" uv run --no-dev python -c "from pathlib import Path; files=sorted(Path('data/processed/full').glob('*.npz')); assert len(files) >= 50; print(f'processed_full_cases={len(files)}')"
""" """
[job] [job]
@ -57,7 +57,7 @@ heartbeat_file = "artifacts/current_run/heartbeat.json"
metrics_file = "artifacts/current_run/metrics.jsonl" metrics_file = "artifacts/current_run/metrics.jsonl"
[artifacts] [artifacts]
mode = "object_store_upload" mode = "rsync"
required = [ required = [
"config.toml", "config.toml",
"metrics.jsonl", "metrics.jsonl",
@ -70,14 +70,10 @@ required = [
"split_manifest.json", "split_manifest.json",
"data_manifest.json", "data_manifest.json",
"normalization.json", "normalization.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json", "run_manifest.json",
"environment_manifest.json",
"artifact_manifest.json", "artifact_manifest.json",
"checksums.txt", "checksums.txt",
"verification_report.json",
] ]
[cleanup] [cleanup]

View file

@ -6,7 +6,7 @@ requires-python = ">=3.11"
dependencies = [ dependencies = [
"huggingface-hub>=0.36.0", "huggingface-hub>=0.36.0",
"numpy>=2.4.0", "numpy>=2.4.0",
"torch>=2.7.1,<2.8.0", "torch>=2.8.0",
"wandb>=0.23.0", "wandb>=0.23.0",
] ]
@ -31,7 +31,3 @@ dev = [
"skypilot[vast]>=0.12.3.post1", "skypilot[vast]>=0.12.3.post1",
"pytest>=9.1.1", "pytest>=9.1.1",
] ]
[tool.pytest.ini_options]
testpaths = ["tests"]

View file

@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
import argparse import argparse
import json
import sys import sys
from airfrans_frontier.paths import DEFAULT_RAW_DATA_DIR, DEFAULT_RAW_MANIFEST_PATH, resolve_path from airfrans_frontier.paths import DEFAULT_RAW_DATA_DIR, DEFAULT_RAW_MANIFEST_PATH, resolve_path
@ -25,50 +24,17 @@ def build_parser() -> argparse.ArgumentParser:
process_raw.add_argument("--force", action="store_true") process_raw.add_argument("--force", action="store_true")
process_raw.set_defaults(command="process-raw") process_raw.set_defaults(command="process-raw")
publish_processed = subparsers.add_parser("publish-processed-hf", help="publish processed .npz data to a Hugging Face dataset repo")
publish_processed.add_argument("--data-root", required=True)
publish_processed.add_argument("--repo-id", required=True)
publish_processed.add_argument("--path-in-repo", default="processed/full")
publish_processed.add_argument("--private", action="store_true")
publish_processed.add_argument("--manifest-out")
publish_processed.set_defaults(command="publish-processed-hf")
prepare_public = subparsers.add_parser(
"prepare-public-hf",
help="download public AirfRANS OF_dataset.zip, process it, and publish processed .npz files to HF",
)
prepare_public.add_argument("--repo-id", required=True)
prepare_public.add_argument("--path-in-repo", default="processed/full")
prepare_public.add_argument("--work-dir", default="artifacts/public_airfrans")
prepare_public.add_argument("--output-dir", default="artifacts/data_cache/airfrans_processed/processed/full")
prepare_public.add_argument("--source-url", default="https://data.isir.upmc.fr/extrality/NeurIPS_2022/OF_dataset.zip")
prepare_public.add_argument("--min-cases", type=int, default=1000)
prepare_public.add_argument("--private", action="store_true")
prepare_public.add_argument("--force", action="store_true")
prepare_public.set_defaults(command="prepare-public-hf")
train = subparsers.add_parser("train", help="train a configured baseline model") train = subparsers.add_parser("train", help="train a configured baseline model")
train.add_argument("config", help="path to a training config TOML file") train.add_argument("config", help="path to a training config TOML file")
train.add_argument("--resume", help="path to checkpoint_latest.pt to resume from") train.add_argument("--resume", help="path to checkpoint_latest.pt to resume from")
train.set_defaults(command="train") train.set_defaults(command="train")
sanity = subparsers.add_parser("model-sanity", help="run toy loss-decrease checks for frontier model families")
sanity.add_argument("--artifact-dir", default="artifacts/model_sanity")
sanity.add_argument("--device", choices=("auto", "cuda", "cpu"), default="auto")
sanity.add_argument("--steps", type=int, default=80)
sanity.add_argument("--families", nargs="*", help="model families to check; defaults to every frontier family")
sanity.set_defaults(command="model-sanity")
return parser return parser
def main(argv: list[str] | None = None) -> int: def main(argv: list[str] | None = None) -> int:
parser = build_parser() parser = build_parser()
args = parser.parse_args(argv) args = parser.parse_args(argv)
if args.command in {"process-raw", "prepare-public-hf", "train", "model-sanity"}:
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
if args.command == "inspect-raw": if args.command == "inspect-raw":
if args.sample_limit < 0: if args.sample_limit < 0:
@ -108,49 +74,10 @@ def main(argv: list[str] | None = None) -> int:
print(f"manifest: {result.manifest_path}") print(f"manifest: {result.manifest_path}")
return 0 return 0
if args.command == "publish-processed-hf":
from airfrans_frontier.training.data_sources import publish_processed_dataset
try:
manifest = publish_processed_dataset(
data_root=resolve_path(args.data_root),
repo_id=args.repo_id,
path_in_repo=args.path_in_repo,
private=args.private,
manifest_out=resolve_path(args.manifest_out) if args.manifest_out else None,
)
except (FileNotFoundError, NotADirectoryError, ValueError, RuntimeError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
print(f"repo_url: {manifest['repo_url']}")
print(f"path_in_repo: {manifest['path_in_repo']}")
print(f"npz_files: {manifest['npz_file_count']}")
return 0
if args.command == "prepare-public-hf":
if args.min_cases <= 0:
print("error: --min-cases must be positive", file=sys.stderr)
return 1
from airfrans_frontier.raw.public import ensure_public_airfrans_processed_hf
try:
report = ensure_public_airfrans_processed_hf(
repo_id=args.repo_id,
path_in_repo=args.path_in_repo,
work_dir=resolve_path(args.work_dir),
output_dir=resolve_path(args.output_dir),
source_url=args.source_url,
min_cases=args.min_cases,
private=args.private,
force=args.force,
)
except (FileNotFoundError, NotADirectoryError, ValueError, RuntimeError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
print(json.dumps(report, indent=2, sort_keys=True))
return 0
if args.command == "train": if args.command == "train":
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
from airfrans_frontier.training.loop import train_from_config_path from airfrans_frontier.training.loop import train_from_config_path
try: try:
@ -163,27 +90,6 @@ def main(argv: list[str] | None = None) -> int:
print(f"final_metrics: {result.run_dir / 'final_metrics.json'}") print(f"final_metrics: {result.run_dir / 'final_metrics.json'}")
return 0 return 0
if args.command == "model-sanity":
if args.steps <= 0:
print("error: --steps must be positive", file=sys.stderr)
return 1
from airfrans_frontier.training.sanity import MODEL_FAMILIES, run_model_sanity
families = tuple(args.families) if args.families else MODEL_FAMILIES
try:
result = run_model_sanity(
artifact_dir=resolve_path(args.artifact_dir),
device_type=args.device,
families=families,
steps=args.steps,
)
except (FileNotFoundError, NotADirectoryError, ValueError, RuntimeError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
print(f"report: {resolve_path(args.artifact_dir) / 'model_sanity_results.json'}")
print(f"families: {len(result['families'])}")
return 0
parser.error(f"unknown command: {args.command}") parser.error(f"unknown command: {args.command}")
return 2 return 2

View file

@ -1,23 +1,6 @@
"""Baseline model definitions.""" """Baseline model definitions."""
from airfrans_frontier.models.film import FourierFiLMMLP from airfrans_frontier.models.film import FourierFiLMMLP
from airfrans_frontier.models.frontier import (
DeepONetBranchTrunk,
LocalPointTransformer,
NeRFCFDMultiRes,
PointContextPerceiver,
RasterFNOUNet,
SirenConditionedINR,
)
from airfrans_frontier.models.mlp import PointwiseMLP from airfrans_frontier.models.mlp import PointwiseMLP
__all__ = [ __all__ = ["FourierFiLMMLP", "PointwiseMLP"]
"DeepONetBranchTrunk",
"FourierFiLMMLP",
"LocalPointTransformer",
"NeRFCFDMultiRes",
"PointContextPerceiver",
"PointwiseMLP",
"RasterFNOUNet",
"SirenConditionedINR",
]

View file

@ -1,404 +0,0 @@
from __future__ import annotations
import math
from collections.abc import Sequence
import torch
from torch import nn
from torch.nn import functional as F
class NeRFCFDMultiRes(nn.Module):
def __init__(
self,
*,
feature_names: Sequence[str],
output_dim: int,
coordinate_features: Sequence[str],
encoding_levels: int,
hidden_width: int,
depth: int,
condition_width: int,
condition_depth: int,
activation: str,
) -> None:
super().__init__()
coordinate_indices = _indices(feature_names, coordinate_features)
condition_indices = _complement_indices(len(feature_names), coordinate_indices)
self.register_buffer("coordinate_indices", torch.tensor(coordinate_indices, dtype=torch.long), persistent=False)
self.register_buffer("condition_indices", torch.tensor(condition_indices, dtype=torch.long), persistent=False)
self.encoding_levels = int(encoding_levels)
encoded_dim = len(coordinate_indices) * (1 + 2 * self.encoding_levels)
condition_input_dim = len(condition_indices) if condition_indices else 1
self.condition_encoder = _mlp(
input_dim=condition_input_dim,
hidden_width=condition_width,
output_dim=condition_width,
depth=condition_depth,
activation=activation,
)
self.decoder = _mlp(
input_dim=encoded_dim + condition_width,
hidden_width=hidden_width,
output_dim=output_dim,
depth=depth,
activation=activation,
activate_output=False,
)
def forward(self, features: torch.Tensor) -> torch.Tensor:
coordinates = features.index_select(dim=1, index=self.coordinate_indices)
condition = _gather_or_zeros(features, self.condition_indices)
encoded = _multires_encode(coordinates, self.encoding_levels)
condition_embedding = self.condition_encoder(condition)
return self.decoder(torch.cat((encoded, condition_embedding), dim=1))
class DeepONetBranchTrunk(nn.Module):
def __init__(
self,
*,
feature_names: Sequence[str],
output_dim: int,
coordinate_features: Sequence[str],
fourier_scales: Sequence[float],
hidden_width: int,
depth: int,
condition_width: int,
condition_depth: int,
activation: str,
) -> None:
super().__init__()
coordinate_indices = _indices(feature_names, coordinate_features)
condition_indices = _complement_indices(len(feature_names), coordinate_indices)
self.register_buffer("coordinate_indices", torch.tensor(coordinate_indices, dtype=torch.long), persistent=False)
self.register_buffer("condition_indices", torch.tensor(condition_indices, dtype=torch.long), persistent=False)
self.register_buffer("fourier_scales", torch.tensor(tuple(float(scale) for scale in fourier_scales), dtype=torch.float32), persistent=False)
trunk_input_dim = len(coordinate_indices) * (1 + 2 * len(fourier_scales))
condition_input_dim = len(condition_indices) if condition_indices else 1
self.branch = _mlp(
input_dim=condition_input_dim,
hidden_width=condition_width,
output_dim=hidden_width,
depth=condition_depth,
activation=activation,
)
self.trunk = _mlp(
input_dim=trunk_input_dim,
hidden_width=hidden_width,
output_dim=hidden_width,
depth=depth,
activation=activation,
)
self.head = nn.Linear(hidden_width, output_dim)
def forward(self, features: torch.Tensor) -> torch.Tensor:
coordinates = features.index_select(dim=1, index=self.coordinate_indices)
condition = _gather_or_zeros(features, self.condition_indices)
trunk = self.trunk(_fourier_features(coordinates, self.fourier_scales))
branch = self.branch(condition)
return self.head(trunk * branch)
class PointContextPerceiver(nn.Module):
def __init__(
self,
*,
input_dim: int,
output_dim: int,
hidden_width: int,
latent_width: int,
context_points: int,
attention_depth: int,
activation: str,
) -> None:
super().__init__()
self.context_tokens = nn.Parameter(torch.empty(context_points, latent_width))
nn.init.normal_(self.context_tokens, std=latent_width ** -0.5)
self.input_projection = nn.Linear(input_dim, latent_width)
self.blocks = nn.ModuleList(
[_PerceiverPointBlock(latent_width=latent_width, hidden_width=hidden_width, activation=activation) for _ in range(attention_depth)]
)
self.head = _mlp(
input_dim=latent_width,
hidden_width=hidden_width,
output_dim=output_dim,
depth=max(1, attention_depth),
activation=activation,
activate_output=False,
)
def forward(self, features: torch.Tensor) -> torch.Tensor:
hidden = self.input_projection(features)
tokens = self.context_tokens.to(dtype=hidden.dtype, device=hidden.device)
for block in self.blocks:
hidden = block(hidden, tokens)
return self.head(hidden)
class LocalPointTransformer(nn.Module):
def __init__(
self,
*,
feature_names: Sequence[str],
output_dim: int,
coordinate_features: Sequence[str],
hidden_width: int,
depth: int,
neighbors: int,
activation: str,
) -> None:
super().__init__()
coordinate_indices = _indices(feature_names, coordinate_features)
self.register_buffer("coordinate_indices", torch.tensor(coordinate_indices, dtype=torch.long), persistent=False)
self.neighbors = int(neighbors)
self.input_projection = nn.Linear(len(feature_names), hidden_width)
self.blocks = nn.ModuleList([_LocalPointBlock(hidden_width=hidden_width, activation=activation) for _ in range(depth)])
self.head = nn.Linear(hidden_width, output_dim)
def forward(self, features: torch.Tensor) -> torch.Tensor:
coordinates = features.index_select(dim=1, index=self.coordinate_indices)
hidden = self.input_projection(features)
for block in self.blocks:
hidden = block(hidden, coordinates, self.neighbors)
return self.head(hidden)
class RasterFNOUNet(nn.Module):
def __init__(
self,
*,
feature_names: Sequence[str],
output_dim: int,
coordinate_features: Sequence[str],
grid_resolution: int,
hidden_width: int,
depth: int,
condition_width: int,
condition_depth: int,
activation: str,
) -> None:
super().__init__()
coordinate_indices = _indices(feature_names, coordinate_features[:2])
condition_indices = _complement_indices(len(feature_names), coordinate_indices)
self.register_buffer("coordinate_indices", torch.tensor(coordinate_indices, dtype=torch.long), persistent=False)
self.register_buffer("condition_indices", torch.tensor(condition_indices, dtype=torch.long), persistent=False)
self.grid_resolution = int(grid_resolution)
self.grid = nn.Parameter(torch.empty(self.grid_resolution, self.grid_resolution, hidden_width))
nn.init.normal_(self.grid, std=hidden_width ** -0.5)
condition_input_dim = len(condition_indices) if condition_indices else 1
self.condition_encoder = _mlp(
input_dim=condition_input_dim,
hidden_width=condition_width,
output_dim=condition_width,
depth=condition_depth,
activation=activation,
)
self.decoder = _mlp(
input_dim=hidden_width + condition_width,
hidden_width=hidden_width,
output_dim=output_dim,
depth=depth,
activation=activation,
activate_output=False,
)
def forward(self, features: torch.Tensor) -> torch.Tensor:
xy = features.index_select(dim=1, index=self.coordinate_indices)
sampled = _sample_grid(self.grid, xy)
condition = self.condition_encoder(_gather_or_zeros(features, self.condition_indices))
return self.decoder(torch.cat((sampled, condition), dim=1))
class SirenConditionedINR(nn.Module):
def __init__(
self,
*,
feature_names: Sequence[str],
output_dim: int,
coordinate_features: Sequence[str],
hidden_width: int,
depth: int,
condition_width: int,
condition_depth: int,
omega0: float,
activation: str,
) -> None:
super().__init__()
coordinate_indices = _indices(feature_names, coordinate_features)
condition_indices = _complement_indices(len(feature_names), coordinate_indices)
self.register_buffer("coordinate_indices", torch.tensor(coordinate_indices, dtype=torch.long), persistent=False)
self.register_buffer("condition_indices", torch.tensor(condition_indices, dtype=torch.long), persistent=False)
condition_input_dim = len(condition_indices) if condition_indices else 1
self.condition_encoder = _mlp(
input_dim=condition_input_dim,
hidden_width=condition_width,
output_dim=condition_width,
depth=condition_depth,
activation=activation,
)
layers: list[nn.Module] = []
input_dim = len(coordinate_indices) + condition_width
for layer_index in range(depth):
layers.append(_SineLayer(input_dim if layer_index == 0 else hidden_width, hidden_width, omega0=omega0, first=layer_index == 0))
self.net = nn.Sequential(*layers)
self.head = nn.Linear(hidden_width, output_dim)
def forward(self, features: torch.Tensor) -> torch.Tensor:
coordinates = features.index_select(dim=1, index=self.coordinate_indices)
condition = self.condition_encoder(_gather_or_zeros(features, self.condition_indices))
hidden = self.net(torch.cat((coordinates, condition), dim=1))
return self.head(hidden)
class _PerceiverPointBlock(nn.Module):
def __init__(self, *, latent_width: int, hidden_width: int, activation: str) -> None:
super().__init__()
self.norm = nn.LayerNorm(latent_width)
self.ffn = _mlp(
input_dim=latent_width,
hidden_width=hidden_width,
output_dim=latent_width,
depth=2,
activation=activation,
activate_output=False,
)
def forward(self, hidden: torch.Tensor, tokens: torch.Tensor) -> torch.Tensor:
scale = hidden.shape[1] ** -0.5
attention = torch.softmax(hidden @ tokens.T * scale, dim=1)
context = attention @ tokens
return hidden + self.ffn(self.norm(hidden + context))
class _LocalPointBlock(nn.Module):
def __init__(self, *, hidden_width: int, activation: str) -> None:
super().__init__()
self.norm = nn.LayerNorm(hidden_width)
self.update = _mlp(
input_dim=hidden_width * 2,
hidden_width=hidden_width,
output_dim=hidden_width,
depth=2,
activation=activation,
activate_output=False,
)
def forward(self, hidden: torch.Tensor, coordinates: torch.Tensor, neighbors: int) -> torch.Tensor:
if hidden.shape[0] <= 1 or neighbors <= 0:
neighborhood = hidden
else:
k = min(neighbors + 1, hidden.shape[0])
distances = torch.cdist(coordinates.float(), coordinates.float())
indices = distances.topk(k=k, largest=False).indices[:, 1:] if k > 1 else distances.topk(k=k, largest=False).indices
gathered = hidden.index_select(dim=0, index=indices.reshape(-1)).reshape(hidden.shape[0], -1, hidden.shape[1])
neighborhood = gathered.mean(dim=1)
return hidden + self.update(torch.cat((self.norm(hidden), neighborhood), dim=1))
class _SineLayer(nn.Module):
def __init__(self, input_dim: int, output_dim: int, *, omega0: float, first: bool) -> None:
super().__init__()
self.linear = nn.Linear(input_dim, output_dim)
self.omega0 = float(omega0)
with torch.no_grad():
bound = 1.0 / input_dim if first else math.sqrt(6.0 / input_dim) / self.omega0
self.linear.weight.uniform_(-bound, bound)
def forward(self, values: torch.Tensor) -> torch.Tensor:
return torch.sin(self.omega0 * self.linear(values))
def _indices(feature_names: Sequence[str], selected_names: Sequence[str]) -> tuple[int, ...]:
indices: list[int] = []
for name in selected_names:
try:
indices.append(tuple(feature_names).index(name))
except ValueError as exc:
raise ValueError(f"Coordinate feature {name!r} is not present in dataset features") from exc
if not indices:
raise ValueError("At least one coordinate feature is required")
return tuple(indices)
def _complement_indices(size: int, excluded: Sequence[int]) -> tuple[int, ...]:
excluded_set = set(excluded)
return tuple(index for index in range(size) if index not in excluded_set)
def _gather_or_zeros(features: torch.Tensor, indices: torch.Tensor) -> torch.Tensor:
if indices.numel() == 0:
return features.new_zeros((features.shape[0], 1))
return features.index_select(dim=1, index=indices)
def _multires_encode(coordinates: torch.Tensor, levels: int) -> torch.Tensor:
if levels <= 0:
return coordinates
pieces = [coordinates]
for level in range(levels):
scale = float(2**level) * math.pi
pieces.append(torch.sin(coordinates * scale))
pieces.append(torch.cos(coordinates * scale))
return torch.cat(pieces, dim=1)
def _fourier_features(coordinates: torch.Tensor, scales: torch.Tensor) -> torch.Tensor:
if scales.numel() == 0:
return coordinates
phases = coordinates.unsqueeze(-1) * scales.to(device=coordinates.device, dtype=coordinates.dtype) * math.pi
return torch.cat((coordinates, torch.sin(phases).flatten(1), torch.cos(phases).flatten(1)), dim=1)
def _sample_grid(grid: torch.Tensor, xy: torch.Tensor) -> torch.Tensor:
resolution = grid.shape[0]
if xy.shape[1] < 2:
raise ValueError("Raster model requires at least x and y coordinate features")
scaled = ((xy[:, :2].clamp(-1.0, 1.0) + 1.0) * 0.5) * float(resolution - 1)
x = scaled[:, 0]
y = scaled[:, 1]
x0 = torch.floor(x).long().clamp(0, resolution - 1)
y0 = torch.floor(y).long().clamp(0, resolution - 1)
x1 = (x0 + 1).clamp(0, resolution - 1)
y1 = (y0 + 1).clamp(0, resolution - 1)
wx = (x - x0.to(x.dtype)).unsqueeze(1)
wy = (y - y0.to(y.dtype)).unsqueeze(1)
g00 = grid[y0, x0]
g10 = grid[y0, x1]
g01 = grid[y1, x0]
g11 = grid[y1, x1]
return (1 - wx) * (1 - wy) * g00 + wx * (1 - wy) * g10 + (1 - wx) * wy * g01 + wx * wy * g11
def _mlp(
*,
input_dim: int,
hidden_width: int,
output_dim: int,
depth: int,
activation: str,
activate_output: bool = True,
) -> nn.Sequential:
layers: list[nn.Module] = []
current_dim = input_dim
for _ in range(max(depth - 1, 0)):
layers.append(nn.Linear(current_dim, hidden_width))
layers.append(_activation(activation))
current_dim = hidden_width
layers.append(nn.Linear(current_dim, output_dim))
if activate_output:
layers.append(_activation(activation))
return nn.Sequential(*layers)
def _activation(name: str) -> nn.Module:
normalized = name.lower()
if normalized == "gelu":
return nn.GELU()
if normalized == "relu":
return nn.ReLU()
if normalized == "silu":
return nn.SiLU()
if normalized == "tanh":
return nn.Tanh()
raise ValueError(f"Unsupported activation: {name}")

View file

@ -74,7 +74,6 @@ def process_raw_dataset(
*, *,
limit: int | None = None, limit: int | None = None,
force: bool = False, force: bool = False,
progress_every: int | None = None,
) -> ProcessingResult: ) -> ProcessingResult:
raw_root = Path(raw_dir).expanduser() raw_root = Path(raw_dir).expanduser()
if not raw_root.is_dir(): if not raw_root.is_dir():
@ -91,27 +90,15 @@ def process_raw_dataset(
records: list[dict[str, object]] = [] records: list[dict[str, object]] = []
started = time.perf_counter() started = time.perf_counter()
total_points = 0 total_points = 0
for index, case_dir in enumerate(case_dirs, start=1): for case_dir in case_dirs:
record, points = process_raw_case_to_npz(case_dir, out_root, force=force) target_path = out_root / f"{case_dir.name}.npz"
total_points += points
records.append(record)
if progress_every is not None and progress_every > 0 and (index % progress_every == 0 or index == len(case_dirs)):
print(f"processed_airfrans_cases={index}/{len(case_dirs)} total_points={total_points}", flush=True)
return write_processing_manifest(out_root, raw_root, records=records, total_points=total_points, started=started)
def process_raw_case_to_npz(case_dir: str | Path, output_dir: str | Path, *, force: bool = False) -> tuple[dict[str, object], int]:
case_path = Path(case_dir).expanduser()
out_root = Path(output_dir).expanduser()
out_root.mkdir(parents=True, exist_ok=True)
target_path = out_root / f"{case_path.name}.npz"
if target_path.exists() and not force: if target_path.exists() and not force:
with np.load(target_path, allow_pickle=False) as npz: with np.load(target_path, allow_pickle=False) as npz:
points = int(npz["features"].shape[0]) points = int(npz["features"].shape[0])
return {"case_id": case_path.name, "path": str(target_path), "points": points, "skipped_existing": True}, points records.append({"case_id": case_dir.name, "path": str(target_path), "points": points, "skipped_existing": True})
total_points += points
metadata, features, targets = process_raw_case(case_path) continue
metadata, features, targets = process_raw_case(case_dir)
_atomic_save_npz( _atomic_save_npz(
target_path, target_path,
features=features, features=features,
@ -121,20 +108,11 @@ def process_raw_case_to_npz(case_dir: str | Path, output_dir: str | Path, *, for
metadata=json.dumps(_metadata_json(metadata), sort_keys=True), metadata=json.dumps(_metadata_json(metadata), sort_keys=True),
) )
points = int(features.shape[0]) points = int(features.shape[0])
return {"case_id": case_path.name, "path": str(target_path), "points": points, "metadata": _metadata_json(metadata)}, points total_points += points
records.append({"case_id": case_dir.name, "path": str(target_path), "points": points, "metadata": _metadata_json(metadata)})
def write_processing_manifest(
output_dir: str | Path,
raw_dir: str | Path,
*,
records: list[dict[str, object]],
total_points: int,
started: float,
) -> ProcessingResult:
out_root = Path(output_dir).expanduser()
manifest = { manifest = {
"raw_dir": str(raw_dir), "raw_dir": str(raw_root),
"output_dir": str(out_root), "output_dir": str(out_root),
"case_count": len(records), "case_count": len(records),
"total_points": total_points, "total_points": total_points,

View file

@ -1,815 +0,0 @@
from __future__ import annotations
import json
import os
import shutil
import struct
import time
import urllib.error
import urllib.request
import urllib.parse
import zlib
import zipfile
from dataclasses import dataclass
from pathlib import Path, PurePosixPath
from typing import Any, Protocol
from airfrans_frontier.training.data_sources import publish_processed_dataset
PUBLIC_OF_DATASET_URL = "https://data.isir.upmc.fr/extrality/NeurIPS_2022/OF_dataset.zip"
DEFAULT_PUBLIC_WORK_DIR = Path("artifacts/public_airfrans")
DEFAULT_PUBLIC_OUTPUT_DIR = Path("artifacts/data_cache/airfrans_processed/processed/full")
_EOCD_SIGNATURE = b"PK\x05\x06"
_ZIP64_EOCD_LOCATOR_SIGNATURE = 0x07064B50
_ZIP64_EOCD_SIGNATURE = 0x06064B50
_CENTRAL_DIRECTORY_SIGNATURE = 0x02014B50
_LOCAL_FILE_HEADER_SIGNATURE = 0x04034B50
_ZIP64_EXTRA_ID = 0x0001
_ZIP64_LIMIT_16 = 0xFFFF
_ZIP64_LIMIT_32 = 0xFFFFFFFF
class RangeReader(Protocol):
size: int
bytes_read: int
def read_range(self, start: int, length: int) -> bytes: ...
@dataclass(frozen=True)
class RemoteZipMember:
filename: str
flag_bits: int
compress_type: int
compress_size: int
file_size: int
header_offset: int
@property
def is_dir(self) -> bool:
return self.filename.endswith("/")
@dataclass(frozen=True)
class StreamingZipProcessingResult:
processing: object
source_bytes: int
ranged_bytes_read: int
class PathRangeReader:
def __init__(self, path: str | Path) -> None:
self.path = Path(path).expanduser()
self.size = self.path.stat().st_size
self.bytes_read = 0
def read_range(self, start: int, length: int) -> bytes:
_validate_range(start, length, self.size)
if length == 0:
return b""
with self.path.open("rb") as handle:
handle.seek(start)
data = handle.read(length)
if len(data) != length:
raise RuntimeError(f"Local range read returned {len(data)} bytes; expected {length}")
self.bytes_read += len(data)
return data
class HttpRangeReader:
def __init__(self, url: str) -> None:
self.url = url
size = _remote_content_length(url)
if size is None:
raise RuntimeError(f"Could not determine remote content length for range streaming: {url}")
self.size = size
self.bytes_read = 0
def read_range(self, start: int, length: int) -> bytes:
_validate_range(start, length, self.size)
if length == 0:
return b""
end = start + length - 1
request = urllib.request.Request(self.url, headers={"Range": f"bytes={start}-{end}"})
try:
with urllib.request.urlopen(request, timeout=60) as response:
status = getattr(response, "status", None)
data = response.read()
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")
raise RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc.code} {body}") from exc
except OSError as exc:
raise RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc}") from exc
if status != 206:
raise RuntimeError(f"Server did not honor HTTP Range for {self.url}: status={status}")
if len(data) != length:
raise RuntimeError(f"HTTP range read returned {len(data)} bytes; expected {length}")
self.bytes_read += len(data)
return data
def ensure_public_airfrans_processed_hf(
*,
repo_id: str,
path_in_repo: str = "processed/full",
work_dir: str | Path = DEFAULT_PUBLIC_WORK_DIR,
output_dir: str | Path = DEFAULT_PUBLIC_OUTPUT_DIR,
source_url: str = PUBLIC_OF_DATASET_URL,
min_cases: int = 1000,
private: bool = False,
force: bool = False,
) -> dict[str, Any]:
if min_cases <= 0:
raise ValueError("min_cases must be positive")
prefix = path_in_repo.strip("/")
started = time.time()
existing = _hf_dataset_status(repo_id=repo_id, path_in_repo=prefix)
if not force and existing["npz_file_count"] >= min_cases and existing["has_manifest"]:
return {
"ok": True,
"phase": "already_published",
"repo_id": repo_id,
"repo_type": "dataset",
"repo_url": f"https://huggingface.co/datasets/{repo_id}",
"path_in_repo": prefix,
"min_cases": min_cases,
"elapsed_seconds": time.time() - started,
**existing,
}
work_root = Path(work_dir).expanduser()
output_root = Path(output_dir).expanduser()
work_root.mkdir(parents=True, exist_ok=True)
output_root.mkdir(parents=True, exist_ok=True)
scratch_root = work_root / "streaming_raw"
print(f"range_stream_process_airfrans_zip source={source_url} output_dir={output_root}", flush=True)
streamed = process_of_dataset_url_streaming(
source_url,
output_root,
scratch_dir=scratch_root,
min_cases=min_cases,
force=force,
progress_every=25,
)
processed = streamed.processing
if processed.case_count < min_cases:
raise RuntimeError(f"Processed only {processed.case_count} cases from public AirfRANS archive; expected at least {min_cases}")
print(f"publish_airfrans_processed_hf repo={repo_id} path_in_repo={prefix}", flush=True)
publish = publish_processed_dataset(
data_root=output_root,
repo_id=repo_id,
path_in_repo=prefix,
private=private,
manifest_out=output_root / "hf_dataset_manifest.json",
)
final = _hf_dataset_status(repo_id=repo_id, path_in_repo=prefix)
if final["npz_file_count"] < min_cases:
raise RuntimeError(f"Published dataset has {final['npz_file_count']} .npz files under {prefix}; expected at least {min_cases}")
if not final["has_manifest"]:
raise RuntimeError(f"Published dataset is missing hf_dataset_manifest.json under {prefix}")
return {
"ok": True,
"phase": "published",
"repo_id": repo_id,
"repo_type": "dataset",
"repo_url": f"https://huggingface.co/datasets/{repo_id}",
"path_in_repo": prefix,
"source_url": source_url,
"streaming": True,
"streaming_mode": "zip_range",
"streaming_scratch_dir": str(scratch_root),
"source_bytes": streamed.source_bytes,
"ranged_bytes_read": streamed.ranged_bytes_read,
"output_dir": str(output_root),
"processed_case_count": processed.case_count,
"processed_total_points": processed.total_points,
"processed_manifest_path": str(processed.manifest_path),
"download": {"url": source_url, "mode": "zip_range", "source_bytes": streamed.source_bytes, "ranged_bytes_read": streamed.ranged_bytes_read},
"publish": publish,
"elapsed_seconds": time.time() - started,
**final,
}
def download_file(url: str, destination: str | Path, *, chunk_size: int = 16 * 1024 * 1024) -> dict[str, Any]:
path = Path(destination).expanduser()
path.parent.mkdir(parents=True, exist_ok=True)
expected_size = _remote_content_length(url)
existing_size = path.stat().st_size if path.exists() else 0
if expected_size is not None and existing_size == expected_size:
return {"url": url, "path": str(path), "bytes": existing_size, "resumed": False, "skipped": True}
headers: dict[str, str] = {}
mode = "wb"
resumed = False
if expected_size is not None and 0 < existing_size < expected_size:
headers["Range"] = f"bytes={existing_size}-"
mode = "ab"
resumed = True
print(
f"download_airfrans_zip url={url} path={path} existing_bytes={existing_size} expected_bytes={expected_size}",
flush=True,
)
request = urllib.request.Request(url, headers=headers)
try:
response = urllib.request.urlopen(request, timeout=60)
except urllib.error.HTTPError as exc:
if exc.code == 416 and expected_size is not None and existing_size >= expected_size:
return {"url": url, "path": str(path), "bytes": existing_size, "resumed": False, "skipped": True}
raise
with response:
if resumed and getattr(response, "status", None) != 206:
mode = "wb"
resumed = False
existing_size = 0
written = existing_size
next_report = ((written // 1_000_000_000) + 1) * 1_000_000_000
with path.open(mode) as handle:
while True:
chunk = response.read(chunk_size)
if not chunk:
break
handle.write(chunk)
written += len(chunk)
if written >= next_report:
print(f"downloaded_airfrans_zip_bytes={written}", flush=True)
next_report += 1_000_000_000
final_size = path.stat().st_size
if expected_size is not None and final_size != expected_size:
raise RuntimeError(f"Downloaded {final_size} bytes from {url}, expected {expected_size}")
return {"url": url, "path": str(path), "bytes": final_size, "resumed": resumed, "skipped": False}
def process_of_dataset_url_streaming(
source_url: str,
output_dir: str | Path,
*,
scratch_dir: str | Path,
min_cases: int = 1000,
force: bool = False,
progress_every: int | None = None,
) -> StreamingZipProcessingResult:
if min_cases <= 0:
raise ValueError("min_cases must be positive")
reader = _range_reader_for(source_url)
members = _read_zip_central_directory(reader)
processing = _process_remote_zip_members(
reader,
members,
output_dir,
scratch_dir=scratch_dir,
raw_dir_label=f"{source_url}!OF_dataset",
min_cases=min_cases,
force=force,
progress_every=progress_every,
)
print(
f"range_stream_airfrans_bytes_read={reader.bytes_read} range_stream_airfrans_source_bytes={reader.size}",
flush=True,
)
return StreamingZipProcessingResult(
processing=processing,
source_bytes=reader.size,
ranged_bytes_read=reader.bytes_read,
)
def _range_reader_for(source_url: str) -> RangeReader:
parsed = urllib.parse.urlparse(source_url)
if parsed.scheme in {"http", "https"}:
return HttpRangeReader(source_url)
if parsed.scheme == "file":
return PathRangeReader(Path(urllib.request.url2pathname(parsed.path)))
if not parsed.scheme:
return PathRangeReader(source_url)
raise RuntimeError(f"Unsupported AirfRANS streaming URL scheme: {parsed.scheme}")
def _read_zip_central_directory(reader: RangeReader) -> list[RemoteZipMember]:
tail_size = min(reader.size, 1024 * 1024)
tail_start = reader.size - tail_size
tail = reader.read_range(tail_start, tail_size)
eocd_index = tail.rfind(_EOCD_SIGNATURE)
if eocd_index < 0:
raise RuntimeError("ZIP end-of-central-directory record not found")
eocd_offset = tail_start + eocd_index
eocd = tail[eocd_index : eocd_index + 22]
if len(eocd) < 22:
raise RuntimeError("Truncated ZIP end-of-central-directory record")
(
_signature,
_disk_number,
_central_disk,
disk_entries,
total_entries,
central_size,
central_offset,
_comment_length,
) = struct.unpack("<IHHHHIIH", eocd)
if (
disk_entries == _ZIP64_LIMIT_16
or total_entries == _ZIP64_LIMIT_16
or central_size == _ZIP64_LIMIT_32
or central_offset == _ZIP64_LIMIT_32
):
total_entries, central_size, central_offset = _read_zip64_central_directory_locator(reader, eocd_offset)
central = reader.read_range(central_offset, central_size)
members = _parse_central_directory(central, expected_entries=total_entries)
print(f"range_stream_airfrans_zip_members={len(members)}", flush=True)
return members
def _read_zip64_central_directory_locator(reader: RangeReader, eocd_offset: int) -> tuple[int, int, int]:
locator_offset = eocd_offset - 20
if locator_offset < 0:
raise RuntimeError("ZIP64 end-of-central-directory locator is missing")
locator = reader.read_range(locator_offset, 20)
signature, _disk_with_record, zip64_eocd_offset, _disk_count = struct.unpack("<IIQI", locator)
if signature != _ZIP64_EOCD_LOCATOR_SIGNATURE:
raise RuntimeError("ZIP64 end-of-central-directory locator has invalid signature")
record = reader.read_range(zip64_eocd_offset, 56)
(
record_signature,
_record_size,
_version_made,
_version_needed,
_disk_number,
_central_disk,
_disk_entries,
total_entries,
central_size,
central_offset,
) = struct.unpack("<IQHHIIQQQQ", record)
if record_signature != _ZIP64_EOCD_SIGNATURE:
raise RuntimeError("ZIP64 end-of-central-directory record has invalid signature")
return int(total_entries), int(central_size), int(central_offset)
def _parse_central_directory(central: bytes, *, expected_entries: int) -> list[RemoteZipMember]:
members: list[RemoteZipMember] = []
offset = 0
while offset < len(central):
if offset + 46 > len(central):
raise RuntimeError("Truncated ZIP central directory entry")
fields = struct.unpack_from("<IHHHHHHIIIHHHHHII", central, offset)
signature = fields[0]
if signature != _CENTRAL_DIRECTORY_SIGNATURE:
raise RuntimeError(f"Invalid ZIP central directory signature at offset {offset}")
flag_bits = fields[3]
compress_type = fields[4]
compress_size = fields[8]
file_size = fields[9]
filename_length = fields[10]
extra_length = fields[11]
comment_length = fields[12]
header_offset = fields[16]
name_start = offset + 46
extra_start = name_start + filename_length
comment_start = extra_start + extra_length
next_offset = comment_start + comment_length
if next_offset > len(central):
raise RuntimeError("Truncated ZIP central directory variable fields")
filename_bytes = central[name_start:extra_start]
encoding = "utf-8" if flag_bits & 0x800 else "cp437"
filename = filename_bytes.decode(encoding, errors="replace")
extra = central[extra_start:comment_start]
file_size, compress_size, header_offset = _apply_zip64_extra(
extra,
file_size=file_size,
compress_size=compress_size,
header_offset=header_offset,
)
members.append(
RemoteZipMember(
filename=filename,
flag_bits=flag_bits,
compress_type=compress_type,
compress_size=compress_size,
file_size=file_size,
header_offset=header_offset,
)
)
offset = next_offset
if expected_entries not in (0, len(members)):
raise RuntimeError(f"ZIP central directory entry count mismatch: parsed={len(members)} expected={expected_entries}")
return members
def _apply_zip64_extra(extra: bytes, *, file_size: int, compress_size: int, header_offset: int) -> tuple[int, int, int]:
values_needed = [
file_size == _ZIP64_LIMIT_32,
compress_size == _ZIP64_LIMIT_32,
header_offset == _ZIP64_LIMIT_32,
]
if not any(values_needed):
return file_size, compress_size, header_offset
offset = 0
while offset + 4 <= len(extra):
header_id, data_size = struct.unpack_from("<HH", extra, offset)
data_start = offset + 4
data_end = data_start + data_size
if data_end > len(extra):
raise RuntimeError("Truncated ZIP extra field")
if header_id == _ZIP64_EXTRA_ID:
cursor = data_start
resolved = [file_size, compress_size, header_offset]
for index, needed in enumerate(values_needed):
if needed:
if cursor + 8 > data_end:
raise RuntimeError("Truncated ZIP64 extra field")
resolved[index] = struct.unpack_from("<Q", extra, cursor)[0]
cursor += 8
return int(resolved[0]), int(resolved[1]), int(resolved[2])
offset = data_end
raise RuntimeError("ZIP64 central directory entry missing ZIP64 extra field")
def _process_remote_zip_members(
reader: RangeReader,
members: list[RemoteZipMember],
output_dir: str | Path,
*,
scratch_dir: str | Path,
raw_dir_label: str,
min_cases: int,
force: bool,
progress_every: int | None,
):
from airfrans_frontier.raw.process import process_raw_case_to_npz, write_processing_manifest
out_root = Path(output_dir).expanduser()
scratch_root = Path(scratch_dir).expanduser()
out_root.mkdir(parents=True, exist_ok=True)
if scratch_root.exists():
shutil.rmtree(scratch_root)
scratch_root.mkdir(parents=True, exist_ok=True)
case_members = _remote_archive_case_members(members)
case_names = sorted(case_members)
if len(case_names) < min_cases:
raise RuntimeError(f"AirfRANS archive has {len(case_names)} cases; expected at least {min_cases}")
print(f"range_stream_airfrans_archive_cases={len(case_names)}", flush=True)
records: list[dict[str, object]] = []
total_points = 0
started = time.perf_counter()
for index, case_name in enumerate(case_names, start=1):
case_dir = scratch_root / case_name
target_path = out_root / f"{case_name}.npz"
if target_path.exists() and not force:
record, points = process_raw_case_to_npz(case_dir, out_root, force=False)
else:
try:
_extract_remote_case_members(reader, case_members[case_name], scratch_root)
record, points = process_raw_case_to_npz(case_dir, out_root, force=force)
finally:
if case_dir.exists():
shutil.rmtree(case_dir, ignore_errors=True)
records.append(record)
total_points += points
if progress_every is not None and progress_every > 0 and (index % progress_every == 0 or index == len(case_names)):
print(f"range_streamed_airfrans_cases={index}/{len(case_names)} total_points={total_points}", flush=True)
try:
scratch_root.rmdir()
except OSError:
pass
return write_processing_manifest(
out_root,
raw_dir_label,
records=records,
total_points=total_points,
started=started,
)
def _remote_archive_case_members(members: list[RemoteZipMember]) -> dict[str, list[tuple[RemoteZipMember, PurePosixPath]]]:
cases: dict[str, list[tuple[RemoteZipMember, PurePosixPath]]] = {}
for member in members:
parsed = _case_member_parts_from_name(member.filename)
if parsed is None:
continue
case_name, relative = parsed
cases.setdefault(case_name, []).append((member, relative))
return cases
def _extract_remote_case_members(
reader: RangeReader,
members: list[tuple[RemoteZipMember, PurePosixPath]],
root: Path,
) -> None:
resolved_root = root.resolve()
for member, relative in members:
target = _safe_relative_target(root, relative, resolved_root=resolved_root)
if member.is_dir:
target.mkdir(parents=True, exist_ok=True)
continue
target.parent.mkdir(parents=True, exist_ok=True)
payload = _read_remote_member_payload(reader, member)
target.write_bytes(payload)
def _read_remote_member_payload(reader: RangeReader, member: RemoteZipMember) -> bytes:
if member.flag_bits & 0x1:
raise RuntimeError(f"Encrypted ZIP member is unsupported: {member.filename}")
local_header = reader.read_range(member.header_offset, 30)
(
signature,
_version_needed,
_flag_bits,
_compress_type,
_mod_time,
_mod_date,
_crc,
_compress_size,
_file_size,
filename_length,
extra_length,
) = struct.unpack("<IHHHHHIIIHH", local_header)
if signature != _LOCAL_FILE_HEADER_SIGNATURE:
raise RuntimeError(f"Invalid local ZIP header for member: {member.filename}")
data_offset = member.header_offset + 30 + filename_length + extra_length
compressed = reader.read_range(data_offset, member.compress_size)
if member.compress_type == 0:
payload = compressed
elif member.compress_type == 8:
decompressor = zlib.decompressobj(-15)
payload = decompressor.decompress(compressed) + decompressor.flush()
else:
raise RuntimeError(f"Unsupported ZIP compression method {member.compress_type} for {member.filename}")
if len(payload) != member.file_size:
raise RuntimeError(f"ZIP member size mismatch for {member.filename}: got {len(payload)} expected {member.file_size}")
return payload
def _validate_range(start: int, length: int, size: int) -> None:
if start < 0 or length < 0 or start + length > size:
raise RuntimeError(f"Invalid range start={start} length={length} size={size}")
def process_of_dataset_archive_streaming(
archive_path: str | Path,
output_dir: str | Path,
*,
scratch_dir: str | Path,
min_cases: int = 1000,
force: bool = False,
progress_every: int | None = None,
):
if min_cases <= 0:
raise ValueError("min_cases must be positive")
archive = Path(archive_path).expanduser()
out_root = Path(output_dir).expanduser()
scratch_root = Path(scratch_dir).expanduser()
out_root.mkdir(parents=True, exist_ok=True)
if scratch_root.exists():
shutil.rmtree(scratch_root)
scratch_root.mkdir(parents=True, exist_ok=True)
from airfrans_frontier.raw.process import process_raw_case_to_npz, write_processing_manifest
records: list[dict[str, object]] = []
total_points = 0
started = time.perf_counter()
with zipfile.ZipFile(archive) as zf:
case_members = _archive_case_members(zf.infolist())
case_names = sorted(case_members)
if len(case_names) < min_cases:
raise RuntimeError(f"AirfRANS archive has {len(case_names)} cases; expected at least {min_cases}")
print(f"stream_airfrans_archive_cases={len(case_names)}", flush=True)
for index, case_name in enumerate(case_names, start=1):
case_dir = scratch_root / case_name
target_path = out_root / f"{case_name}.npz"
if target_path.exists() and not force:
record, points = process_raw_case_to_npz(case_dir, out_root, force=False)
else:
try:
_extract_case_members(zf, case_members[case_name], scratch_root)
record, points = process_raw_case_to_npz(case_dir, out_root, force=force)
finally:
if case_dir.exists():
shutil.rmtree(case_dir, ignore_errors=True)
records.append(record)
total_points += points
if progress_every is not None and progress_every > 0 and (index % progress_every == 0 or index == len(case_names)):
print(f"streamed_airfrans_cases={index}/{len(case_names)} total_points={total_points}", flush=True)
try:
scratch_root.rmdir()
except OSError:
pass
return write_processing_manifest(
out_root,
f"{archive}!OF_dataset",
records=records,
total_points=total_points,
started=started,
)
def extract_of_dataset(archive_path: str | Path, extract_root: str | Path, *, min_cases: int = 1000) -> Path:
archive = Path(archive_path).expanduser()
root = Path(extract_root).expanduser()
root.mkdir(parents=True, exist_ok=True)
existing = _find_of_dataset_root(root)
if existing is not None and _case_count(existing) >= min_cases:
return existing
print(f"extract_airfrans_zip archive={archive} root={root}", flush=True)
with zipfile.ZipFile(archive) as zf:
members = zf.infolist()
_require_extract_space(root, members)
for index, member in enumerate(members, start=1):
_safe_extract_member(zf, member, root)
if index % 1000 == 0 or index == len(members):
print(f"extracted_airfrans_members={index}/{len(members)}", flush=True)
found = _find_of_dataset_root(root)
if found is None:
raise RuntimeError(f"OF_dataset directory not found after extracting {archive}")
case_count = _case_count(found)
if case_count < min_cases:
raise RuntimeError(f"Extracted AirfRANS OF_dataset has {case_count} cases; expected at least {min_cases}")
return found
def _hf_dataset_status(*, repo_id: str, path_in_repo: str) -> dict[str, Any]:
try:
from huggingface_hub import HfApi
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required for AirfRANS public data preparation") from exc
token = _optional_secret("HF_TOKEN")
api = HfApi(token=token)
try:
files = api.list_repo_files(repo_id=repo_id, repo_type="dataset")
except Exception:
files = []
prefix = path_in_repo.strip("/")
base = f"{prefix}/" if prefix else ""
npz_count = sum(1 for item in files if item.startswith(base) and item.endswith(".npz"))
has_manifest = any(item == f"{base}hf_dataset_manifest.json" for item in files)
return {
"file_count": len(files),
"npz_file_count": npz_count,
"has_manifest": has_manifest,
}
def _remote_content_length(url: str) -> int | None:
request = urllib.request.Request(url, method="HEAD")
try:
with urllib.request.urlopen(request, timeout=60) as response:
raw = response.headers.get("Content-Length")
except Exception:
return None
if raw is None:
return None
try:
return int(raw)
except ValueError:
return None
def _safe_extract_member(zf: zipfile.ZipFile, member: zipfile.ZipInfo, root: Path) -> None:
target = _safe_member_target(member, root)
if member.is_dir():
target.mkdir(parents=True, exist_ok=True)
return
target.parent.mkdir(parents=True, exist_ok=True)
with zf.open(member) as source, target.open("wb") as destination:
shutil.copyfileobj(source, destination, length=16 * 1024 * 1024)
def _require_extract_space(root: Path, members: list[zipfile.ZipInfo]) -> None:
total_uncompressed_bytes = 0
remaining_uncompressed_bytes = 0
resolved_root = root.resolve()
for member in members:
if member.is_dir():
continue
total_uncompressed_bytes += member.file_size
target = _safe_member_target(member, root, resolved_root=resolved_root)
try:
existing_size = target.stat().st_size
except OSError:
existing_size = None
if existing_size == member.file_size:
continue
remaining_uncompressed_bytes += member.file_size
margin_bytes = max(1024**3, remaining_uncompressed_bytes // 20) if remaining_uncompressed_bytes else 0
required_free_bytes = remaining_uncompressed_bytes + margin_bytes
usage = shutil.disk_usage(root)
print(
"airfrans_extract_total_uncompressed_bytes="
f"{total_uncompressed_bytes} airfrans_extract_remaining_uncompressed_bytes={remaining_uncompressed_bytes} "
f"airfrans_extract_free_disk_bytes={usage.free} airfrans_extract_required_free_bytes={required_free_bytes}",
flush=True,
)
if usage.free < required_free_bytes:
raise RuntimeError(
"Insufficient free disk for AirfRANS extraction: "
f"free={usage.free} required={required_free_bytes} remaining_uncompressed={remaining_uncompressed_bytes}; "
"provision more disk or use a streaming/incremental extraction pipeline"
)
def _safe_member_target(member: zipfile.ZipInfo, root: Path, *, resolved_root: Path | None = None) -> Path:
return _safe_relative_target(root, PurePosixPath(member.filename), resolved_root=resolved_root)
def _safe_relative_target(root: Path, relative: PurePosixPath, *, resolved_root: Path | None = None) -> Path:
target = root.joinpath(*relative.parts)
actual_root = resolved_root or root.resolve()
resolved_target = target.resolve()
if actual_root != resolved_target and actual_root not in resolved_target.parents:
raise RuntimeError(f"Unsafe path in AirfRANS archive: {relative}")
return target
def _archive_case_members(members: list[zipfile.ZipInfo]) -> dict[str, list[tuple[zipfile.ZipInfo, PurePosixPath]]]:
cases: dict[str, list[tuple[zipfile.ZipInfo, PurePosixPath]]] = {}
for member in members:
parsed = _case_member_parts(member)
if parsed is None:
continue
case_name, relative = parsed
cases.setdefault(case_name, []).append((member, relative))
return cases
def _case_member_parts(member: zipfile.ZipInfo) -> tuple[str, PurePosixPath] | None:
return _case_member_parts_from_name(member.filename)
def _case_member_parts_from_name(filename: str) -> tuple[str, PurePosixPath] | None:
parts = PurePosixPath(filename).parts
if any(part == ".." for part in parts):
raise RuntimeError(f"Unsafe path in AirfRANS archive: {filename}")
for index, part in enumerate(parts):
if part.startswith("airFoil2D_"):
return part, PurePosixPath(*parts[index:])
return None
def _extract_case_members(
zf: zipfile.ZipFile,
members: list[tuple[zipfile.ZipInfo, PurePosixPath]],
root: Path,
) -> None:
resolved_root = root.resolve()
for member, relative in members:
target = _safe_relative_target(root, relative, resolved_root=resolved_root)
if member.is_dir():
target.mkdir(parents=True, exist_ok=True)
continue
target.parent.mkdir(parents=True, exist_ok=True)
with zf.open(member) as source, target.open("wb") as destination:
shutil.copyfileobj(source, destination, length=16 * 1024 * 1024)
def _find_of_dataset_root(root: Path) -> Path | None:
direct = root / "OF_dataset"
if direct.is_dir():
return direct
for candidate in root.glob("*/OF_dataset"):
if candidate.is_dir():
return candidate
if _case_count(root) > 0:
return root
return None
def _case_count(root: Path) -> int:
return sum(1 for path in root.iterdir() if path.is_dir() and path.name.startswith("airFoil2D_")) if root.is_dir() else 0
def _optional_secret(name: str) -> str | None:
value = os.environ.get(name)
if value:
return value
for path in (Path(".env") / name, Path(".env") / f"{name}.txt"):
if path.is_file():
text = path.read_text().strip()
if text:
return text
return None
def _remove_file_best_effort(path: Path) -> bool:
try:
path.unlink()
return True
except FileNotFoundError:
return False
except OSError as exc:
print(f"warning: could not remove {path}: {exc}", flush=True)
return False
def write_json_report(path: str | Path, payload: dict[str, Any]) -> None:
report_path = Path(path).expanduser()
report_path.parent.mkdir(parents=True, exist_ok=True)
report_path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")

View file

@ -5,6 +5,7 @@ import json
from pathlib import Path from pathlib import Path
from typing import Any, Iterable from typing import Any, Iterable
import torch
BASE_REQUIRED = ( BASE_REQUIRED = (
"config.toml", "config.toml",
@ -24,7 +25,6 @@ def verify_artifacts(
required: Iterable[str] = DEFAULT_REQUIRED, required: Iterable[str] = DEFAULT_REQUIRED,
*, *,
require_terminal: bool = True, require_terminal: bool = True,
verify_hf_remote: bool = False,
) -> dict[str, Any]: ) -> dict[str, Any]:
root = Path(artifact_dir) root = Path(artifact_dir)
if not root.exists(): if not root.exists():
@ -50,11 +50,6 @@ def verify_artifacts(
if missing_failure: if missing_failure:
raise ValueError(f"Failed artifact directory missing files: {', '.join(missing_failure)}") raise ValueError(f"Failed artifact directory missing files: {', '.join(missing_failure)}")
prior_checks = _prior_verification_checks(root / "verification_report.json")
checks: dict[str, Any] = {
"required_files": {name: True for name in required_names},
"terminal_artifact": "final_metrics.json" if has_final else "failure_report.json" if has_failure else None,
}
for json_name in ( for json_name in (
"latest_metrics.json", "latest_metrics.json",
"heartbeat.json", "heartbeat.json",
@ -64,38 +59,20 @@ def verify_artifacts(
"split_manifest.json", "split_manifest.json",
"data_manifest.json", "data_manifest.json",
"normalization.json", "normalization.json",
"environment_manifest.json",
"evaluation_protocol.json",
"artifact_manifest.json",
"hf_upload_manifest.json",
"artifact_collection_report.json",
"disk_telemetry.json",
"verification_report.json",
): ):
path = root / json_name path = root / json_name
if path.is_file(): if path.is_file():
_validate_json(path) _validate_json(path)
checks[f"json:{json_name}"] = True
if (root / "metrics.jsonl").is_file():
_validate_jsonl(root / "metrics.jsonl") _validate_jsonl(root / "metrics.jsonl")
checks["jsonl:metrics.jsonl"] = True
for checkpoint_name in ("checkpoint_latest.pt", "checkpoint_best.pt", "checkpoint_final.pt"): for checkpoint_name in ("checkpoint_latest.pt", "checkpoint_best.pt", "checkpoint_final.pt"):
path = root / checkpoint_name path = root / checkpoint_name
if path.is_file(): if path.is_file():
checkpoint_check = f"checkpoint:{checkpoint_name}"
if prior_checks.get(checkpoint_check) is not True:
_validate_checkpoint_metadata(path) _validate_checkpoint_metadata(path)
checks[checkpoint_check] = True
if (root / "hf_upload_manifest.json").is_file():
checks["hf_upload_manifest.json"] = _validate_hf_upload_manifest(root / "hf_upload_manifest.json")
if verify_hf_remote:
checks["hf_remote_paths"] = _verify_hf_remote_paths(root / "hf_upload_manifest.json")
files = sorted( files = sorted(
path path
for path in root.rglob("*") for path in root.rglob("*")
if not path.is_symlink() and path.is_file() and path.name not in {"artifact_manifest.json", "checksums.txt", "verification_report.json"} if not path.is_symlink() and path.is_file() and path.name not in {"artifact_manifest.json", "checksums.txt"}
) )
manifest = { manifest = {
"artifact_dir": str(root), "artifact_dir": str(root),
@ -113,17 +90,6 @@ def verify_artifacts(
(root / "checksums.txt").write_text( (root / "checksums.txt").write_text(
"".join(f"{item['sha256']} {item['path']}\n" for item in manifest["files"]) "".join(f"{item['sha256']} {item['path']}\n" for item in manifest["files"])
) )
checks["artifact_manifest.json"] = _validate_artifact_manifest(root / "artifact_manifest.json", root)
checks["checksums.txt"] = _validate_checksums(root / "checksums.txt", root)
report = {
"ok": True,
"artifact_dir": str(root),
"required": list(required_names),
"checked_at": __import__("time").time(),
"checks": checks,
"manifest_file_count": manifest["file_count"],
}
(root / "verification_report.json").write_text(json.dumps(report, indent=2, sort_keys=True) + "\n")
return manifest return manifest
@ -134,16 +100,6 @@ def sha256_file(path: Path) -> str:
digest.update(chunk) digest.update(chunk)
return digest.hexdigest() return digest.hexdigest()
def _prior_verification_checks(path: Path) -> dict[str, Any]:
if not path.is_file():
return {}
try:
report = json.loads(path.read_text())
except json.JSONDecodeError:
return {}
checks = report.get("checks") if isinstance(report, dict) else None
return dict(checks) if isinstance(checks, dict) else {}
def _validate_json(path: Path) -> None: def _validate_json(path: Path) -> None:
try: try:
@ -165,7 +121,6 @@ def _validate_jsonl(path: Path) -> None:
def _validate_checkpoint_metadata(path: Path) -> None: def _validate_checkpoint_metadata(path: Path) -> None:
import torch
try: try:
checkpoint = torch.load(path, map_location="cpu", weights_only=False) checkpoint = torch.load(path, map_location="cpu", weights_only=False)
except Exception as exc: except Exception as exc:
@ -176,82 +131,3 @@ def _validate_checkpoint_metadata(path: Path) -> None:
missing = [name for name in required if name not in checkpoint] missing = [name for name in required if name not in checkpoint]
if missing: if missing:
raise ValueError(f"Checkpoint artifact {path} missing keys: {', '.join(missing)}") raise ValueError(f"Checkpoint artifact {path} missing keys: {', '.join(missing)}")
def _validate_checksums(path: Path, root: Path) -> dict[str, int]:
checked = 0
for line_number, raw_line in enumerate(path.read_text().splitlines(), start=1):
if not raw_line.strip():
continue
try:
expected, relative = raw_line.split(" ", 1)
except ValueError as exc:
raise ValueError(f"Invalid checksum line {path}:{line_number}") from exc
target = root / relative
if not target.is_file():
raise ValueError(f"Checksum references missing artifact: {relative}")
actual = sha256_file(target)
if actual != expected:
raise ValueError(f"Checksum mismatch for artifact: {relative}")
checked += 1
return {"checked": checked}
def _validate_artifact_manifest(path: Path, root: Path) -> dict[str, int]:
data = json.loads(path.read_text())
if not isinstance(data, dict):
raise ValueError(f"Artifact manifest is not a mapping: {path}")
files = data.get("files")
if not isinstance(files, list):
raise ValueError(f"Artifact manifest missing files list: {path}")
checked = 0
for item in files:
if not isinstance(item, dict):
raise ValueError(f"Artifact manifest file entry is not a mapping: {path}")
relative = item.get("path")
expected = item.get("sha256")
if not isinstance(relative, str) or not isinstance(expected, str):
raise ValueError(f"Artifact manifest file entry missing path or sha256: {path}")
target = root / relative
if not target.is_file():
raise ValueError(f"Artifact manifest references missing artifact: {relative}")
if sha256_file(target) != expected:
raise ValueError(f"Artifact manifest checksum mismatch: {relative}")
checked += 1
return {"checked": checked}
def _validate_hf_upload_manifest(path: Path) -> dict[str, Any]:
data = json.loads(path.read_text())
if not isinstance(data, dict):
raise ValueError(f"HF upload manifest is not a mapping: {path}")
enabled = data.get("enabled")
if enabled is False:
return {"enabled": False, "uploaded_paths": 0}
uploaded_paths = data.get("uploaded_paths")
if uploaded_paths is None:
uploaded_paths = []
if not isinstance(uploaded_paths, list) or not all(isinstance(item, str) for item in uploaded_paths):
raise ValueError(f"HF upload manifest uploaded_paths must be a list of strings: {path}")
return {"enabled": bool(enabled), "uploaded_paths": len(uploaded_paths)}
def _verify_hf_remote_paths(path: Path) -> dict[str, Any]:
data = json.loads(path.read_text())
if not data.get("enabled"):
return {"enabled": False}
repo_id = data.get("repo_id")
repo_type = data.get("repo_type", "model")
uploaded_paths = data.get("uploaded_paths", [])
if not isinstance(repo_id, str) or not isinstance(uploaded_paths, list):
raise ValueError(f"HF upload manifest cannot be remote-verified: {path}")
try:
from huggingface_hub import HfApi
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required for remote HF artifact verification") from exc
api = HfApi()
remote_files = set(api.list_repo_files(repo_id=repo_id, repo_type=repo_type))
missing = [item for item in uploaded_paths if item not in remote_files]
if missing:
raise ValueError(f"HF repo is missing uploaded artifact paths: {', '.join(missing)}")
return {"enabled": True, "checked": len(uploaded_paths)}

View file

@ -1,166 +0,0 @@
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
import time
from typing import Any
_TERMINAL_INSTANCE_STATUSES = {
"deleted",
"destroyed",
"exited",
"offline",
"stopped",
"stopping",
"terminated",
}
def reconcile_cleanup(
*,
sky_state: Any,
vast_instances: Sequence[Mapping[str, Any]],
known_run_ids: Sequence[str] = (),
destroy_orphans: bool = False,
destroy_instance: Callable[[int], Any] | None = None,
now: float | None = None,
) -> dict[str, Any]:
"""Reconcile Sky's view with Vast API ground truth and report cleanup actions.
Vast instances are treated as the paid-resource ground truth. Destruction is
opt-in so this can be used as a non-launch-blocking inspection command.
"""
checked_at = time.time() if now is None else float(now)
sky_refs = _extract_sky_refs(sky_state)
known_runs = tuple(known_run_ids)
records: list[dict[str, Any]] = []
for instance in vast_instances:
instance_id = _instance_id(instance)
status = _status(instance)
associated_run_id = _associated_run_id(instance, known_runs)
sky_knows = _sky_knows_instance(sky_refs, instance_id=instance_id, run_id=associated_run_id)
live = _is_live_status(status)
unexpected_live = bool(live and not sky_knows)
action = "none"
result = "not_needed"
error = None
if unexpected_live:
action = "destroy_orphan" if destroy_orphans else "report_orphan"
result = "not_attempted"
if destroy_orphans:
if destroy_instance is None:
result = "skipped_no_destroy_function"
elif instance_id is None:
result = "skipped_missing_instance_id"
else:
try:
destroy_instance(int(instance_id))
except Exception as exc: # pragma: no cover - exercised by callers with fakes.
result = "failed"
error = str(exc)
else:
result = "destroy_requested"
records.append(
{
"vast_instance_id": instance_id,
"host_id": _first_present(instance, "host_id", "machine_id"),
"gpu_type": _first_present(instance, "gpu_name", "gpu", "gpu_type"),
"gpu_count": _first_present(instance, "num_gpus", "gpu_count", "gpus"),
"status": status,
"associated_run_id": associated_run_id,
"hourly_cost": _first_present(instance, "dph_total", "hourly_cost", "cost_per_hour"),
"sky_known": sky_knows,
"live": live,
"unexpected_live": unexpected_live,
"cleanup_action_attempted": action,
"cleanup_result": result,
"error": error,
}
)
return {
"schema_version": 1,
"checked_at": checked_at,
"sky_instance_ids": sorted(sky_refs["instance_ids"]),
"sky_run_ids": sorted(sky_refs["run_ids"]),
"destroy_orphans": destroy_orphans,
"unexpected_live_count": sum(1 for record in records if record["unexpected_live"]),
"destroy_requested_count": sum(1 for record in records if record["cleanup_result"] == "destroy_requested"),
"instances": records,
}
def _extract_sky_refs(value: Any) -> dict[str, set[str]]:
refs = {"instance_ids": set(), "run_ids": set()}
_walk_sky(value, refs)
return refs
def _walk_sky(value: Any, refs: dict[str, set[str]]) -> None:
if isinstance(value, Mapping):
for key, item in value.items():
key_text = str(key).lower()
if key_text in {"id", "instance_id", "vast_instance_id"}:
_add_ref(refs["instance_ids"], item)
elif key_text in {"name", "cluster", "cluster_name", "run_id", "label"}:
_add_ref(refs["run_ids"], item)
_walk_sky(item, refs)
elif isinstance(value, (list, tuple)):
for item in value:
_walk_sky(item, refs)
def _add_ref(target: set[str], value: Any) -> None:
if isinstance(value, bool) or value is None:
return
if isinstance(value, (int, float, str)):
text = str(int(value)) if isinstance(value, float) and value.is_integer() else str(value)
if text:
target.add(text)
def _sky_knows_instance(refs: Mapping[str, set[str]], *, instance_id: int | None, run_id: str | None) -> bool:
if instance_id is not None and str(instance_id) in refs["instance_ids"]:
return True
if run_id is not None and run_id in refs["run_ids"]:
return True
return False
def _instance_id(instance: Mapping[str, Any]) -> int | None:
value = _first_present(instance, "id", "instance_id", "vast_instance_id")
if isinstance(value, bool) or value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def _status(instance: Mapping[str, Any]) -> str | None:
value = _first_present(instance, "actual_status", "status", "state")
return str(value) if value is not None else None
def _is_live_status(status: str | None) -> bool:
if status is None:
return True
return status.lower() not in _TERMINAL_INSTANCE_STATUSES
def _associated_run_id(instance: Mapping[str, Any], known_run_ids: Sequence[str]) -> str | None:
for key in ("run_id", "label", "name", "cluster_name"):
value = instance.get(key)
if isinstance(value, str) and value:
if value in known_run_ids:
return value
for run_id in known_run_ids:
if run_id and run_id in value:
return run_id
return None
def _first_present(instance: Mapping[str, Any], *keys: str) -> Any:
for key in keys:
if key in instance and instance[key] is not None:
return instance[key]
return None

View file

@ -12,15 +12,11 @@ from pathlib import Path
from typing import Any from typing import Any
from airfrans_frontier.remote.artifacts import verify_artifacts from airfrans_frontier.remote.artifacts import verify_artifacts
from airfrans_frontier.remote.cleanup import reconcile_cleanup
from airfrans_frontier.remote.collection import ARTIFACT_COLLECTION_REPORT, collect_artifact_paths, required_collection_failures
from airfrans_frontier.remote.config import RemoteRunConfig, load_remote_run_config from airfrans_frontier.remote.config import RemoteRunConfig, load_remote_run_config
from airfrans_frontier.remote.launch_group import LaunchGroupScheduler
from airfrans_frontier.remote.selection import DEFAULT_SELECTION_MAX_AGE_SECONDS, load_selection_manifest
from airfrans_frontier.remote.skypilot import render_skypilot_yaml, write_skyignore from airfrans_frontier.remote.skypilot import render_skypilot_yaml, write_skyignore
from airfrans_frontier.remote.skypilot_patch import apply_patch, patch_status, require_patch from airfrans_frontier.remote.skypilot_patch import apply_patch, patch_status, require_patch
from airfrans_frontier.remote.smoke import run_hf_upload_smoke, run_smoke_training, run_wandb_smoke from airfrans_frontier.remote.smoke import run_hf_upload_smoke, run_smoke_training, run_wandb_smoke
from airfrans_frontier.remote.vast import SelectionResult, destroy_instance, list_instances, select_offer, summarize_instances from airfrans_frontier.remote.vast import SelectionResult, select_offer
def build_parser() -> argparse.ArgumentParser: def build_parser() -> argparse.ArgumentParser:
@ -31,38 +27,14 @@ def build_parser() -> argparse.ArgumentParser:
doctor.add_argument("--apply-skypilot-patch", action="store_true") doctor.add_argument("--apply-skypilot-patch", action="store_true")
doctor.set_defaults(command="doctor") doctor.set_defaults(command="doctor")
vast_instances = subparsers.add_parser("vast-instances", help="list Vast.ai instances using the Vast API")
vast_instances.add_argument("--base-url", default="https://cloud.vast.ai")
vast_instances.add_argument("--out")
vast_instances.set_defaults(command="vast-instances")
cleanup = subparsers.add_parser("cleanup-reconcile", help="reconcile Sky status against Vast API ground truth")
cleanup.add_argument("--sky-status-json", help="local Sky status JSON; omit to call sky status")
cleanup.add_argument("--vast-instances-json", help="local Vast instances JSON; omit to call Vast API")
cleanup.add_argument("--base-url", default="https://cloud.vast.ai")
cleanup.add_argument("--destroy-orphans", action="store_true", help="request Vast destruction for live instances missing from Sky")
cleanup.add_argument("--out")
cleanup.set_defaults(command="cleanup-reconcile")
select = subparsers.add_parser("select", help="select a Vast.ai offer from a remote config") select = subparsers.add_parser("select", help="select a Vast.ai offer from a remote config")
select.add_argument("config") select.add_argument("config")
select.add_argument("--out") select.add_argument("--out")
select.set_defaults(command="select") select.set_defaults(command="select")
launch_group = subparsers.add_parser("launch-group-plan", help="write local launch-group state without provisioning")
launch_group.add_argument("configs", nargs="+")
launch_group.add_argument("--max-active", type=int, default=4)
launch_group.add_argument("--max-fragile", type=int, default=1)
launch_group.add_argument("--state")
launch_group.add_argument("--group-id")
launch_group.add_argument("--allow-duplicate-hosts", action="store_true")
launch_group.set_defaults(command="launch-group-plan")
render = subparsers.add_parser("render", help="render patched SkyPilot YAML") render = subparsers.add_parser("render", help="render patched SkyPilot YAML")
render.add_argument("config") render.add_argument("config")
render.add_argument("--selection", required=True) render.add_argument("--selection", required=True)
render.add_argument("--selection-max-age-seconds", type=float, default=DEFAULT_SELECTION_MAX_AGE_SECONDS)
render.add_argument("--allow-stale-selection", action="store_true")
render.add_argument("--run-id", required=True) render.add_argument("--run-id", required=True)
render.add_argument("--out") render.add_argument("--out")
render.set_defaults(command="render") render.set_defaults(command="render")
@ -107,57 +79,16 @@ def main(argv: list[str] | None = None) -> int:
parser = build_parser() parser = build_parser()
args = parser.parse_args(argv) args = parser.parse_args(argv)
try: try:
if args.command in {"doctor", "run"}:
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
if args.command == "doctor": if args.command == "doctor":
return _doctor(apply=args.apply_skypilot_patch) return _doctor(apply=args.apply_skypilot_patch)
if args.command == "vast-instances":
api_key = os.environ.get("VAST_API_KEY")
if not api_key:
raise RuntimeError("VAST_API_KEY is required to list Vast instances")
instances = list_instances(base_url=args.base_url, api_key=api_key)
_emit_json({"instance_count": len(instances), "instances": summarize_instances(instances)}, args.out)
return 0
if args.command == "cleanup-reconcile":
sky_state = _load_json_file(Path(args.sky_status_json)) if args.sky_status_json else _load_sky_status()
if args.vast_instances_json:
vast_payload = _load_json_file(Path(args.vast_instances_json))
instances = _instances_from_json_payload(vast_payload)
else:
api_key = os.environ.get("VAST_API_KEY")
if not api_key:
raise RuntimeError("VAST_API_KEY is required to reconcile live Vast instances")
instances = list_instances(base_url=args.base_url, api_key=api_key)
destroy = None
if args.destroy_orphans:
api_key = os.environ.get("VAST_API_KEY")
if not api_key:
raise RuntimeError("VAST_API_KEY is required to destroy Vast orphan instances")
destroy = lambda instance_id: destroy_instance(base_url=args.base_url, api_key=api_key, instance_id=instance_id)
report = reconcile_cleanup(sky_state=sky_state, vast_instances=instances, destroy_orphans=args.destroy_orphans, destroy_instance=destroy)
_emit_json(report, args.out)
return 0
if args.command == "select": if args.command == "select":
config = load_remote_run_config(args.config) config = load_remote_run_config(args.config)
result = select_offer(config) result = select_offer(config)
_emit_json(result.to_manifest(), args.out) _emit_json(result.to_manifest(), args.out)
return 0 return 0
if args.command == "launch-group-plan":
scheduler = LaunchGroupScheduler(
args.configs,
max_active=args.max_active,
max_fragile=args.max_fragile,
state_path=args.state,
group_id=args.group_id,
allow_duplicate_hosts=args.allow_duplicate_hosts,
)
_emit_json(scheduler.to_payload(), None)
return 0
if args.command == "render": if args.command == "render":
config = load_remote_run_config(args.config) config = load_remote_run_config(args.config)
selection = _selection_from_manifest(Path(args.selection), max_age_seconds=args.selection_max_age_seconds, allow_stale=args.allow_stale_selection) selection = _selection_from_manifest(Path(args.selection))
text = render_skypilot_yaml(config, selection, run_id=args.run_id) text = render_skypilot_yaml(config, selection, run_id=args.run_id)
if args.out: if args.out:
Path(args.out).write_text(text) Path(args.out).write_text(text)
@ -165,9 +96,6 @@ def main(argv: list[str] | None = None) -> int:
print(text) print(text)
return 0 return 0
if args.command == "verify-artifacts": if args.command == "verify-artifacts":
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
manifest = verify_artifacts(args.artifact_dir) manifest = verify_artifacts(args.artifact_dir)
print(json.dumps({"status": "ok", "file_count": manifest["file_count"]}, sort_keys=True)) print(json.dumps({"status": "ok", "file_count": manifest["file_count"]}, sort_keys=True))
return 0 return 0
@ -211,7 +139,7 @@ def _doctor(*, apply: bool) -> int:
problems: list[str] = [] problems: list[str] = []
if not os.environ.get("VAST_API_KEY"): if not os.environ.get("VAST_API_KEY"):
problems.append("VAST_API_KEY is not set") problems.append("VAST_API_KEY is not set")
sky = shutil.which("sky", path=_subprocess_env().get("PATH")) sky = shutil.which("sky")
if not sky: if not sky:
problems.append("sky executable not found on PATH") problems.append("sky executable not found on PATH")
if apply: if apply:
@ -245,21 +173,6 @@ def _run(config_path: str | Path, *, dry_run: bool, skip_down: bool) -> int:
local_run_dir = config.run.local_artifact_dir / run_id local_run_dir = config.run.local_artifact_dir / run_id
local_run_dir.mkdir(parents=True, exist_ok=False) local_run_dir.mkdir(parents=True, exist_ok=False)
state_path = local_run_dir / "orchestrator_state.json" state_path = local_run_dir / "orchestrator_state.json"
timeline_path = local_run_dir / "startup_timeline.jsonl"
submitted_at = time.time()
def timeline(phase: str, event: str, **extra: Any) -> None:
record = {
"run_id": run_id,
"ts": time.time(),
"elapsed_since_submit_seconds": time.time() - submitted_at,
"phase": phase,
"event": event,
**extra,
}
with timeline_path.open("a", encoding="utf-8") as handle:
handle.write(json.dumps(record, sort_keys=True) + "\n")
def state(phase: str, **extra: Any) -> None: def state(phase: str, **extra: Any) -> None:
payload = { payload = {
@ -270,7 +183,7 @@ def _run(config_path: str | Path, *, dry_run: bool, skip_down: bool) -> int:
**extra, **extra,
} }
state_path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n") state_path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
timeline("orchestrator", phase.lower(), orchestrator_phase=phase, **extra)
state("SELECTING_OFFER") state("SELECTING_OFFER")
selection = select_offer(config) selection = select_offer(config)
selection_path = local_run_dir / "selection_manifest.json" selection_path = local_run_dir / "selection_manifest.json"
@ -311,7 +224,6 @@ def _run(config_path: str | Path, *, dry_run: bool, skip_down: bool) -> int:
attempt=attempt, attempt=attempt,
resume_checkpoint=str(resume_checkpoint) if resume_checkpoint is not None else None, resume_checkpoint=str(resume_checkpoint) if resume_checkpoint is not None else None,
) )
timeline("sky_launch", "started", attempt=attempt, selected_offer_id=selection.selected_offer_id)
return_code = _run_sky_with_periodic_collection( return_code = _run_sky_with_periodic_collection(
cluster=run_id, cluster=run_id,
sky_yaml_path=sky_yaml_path, sky_yaml_path=sky_yaml_path,
@ -319,7 +231,6 @@ def _run(config_path: str | Path, *, dry_run: bool, skip_down: bool) -> int:
local_run_dir=local_run_dir, local_run_dir=local_run_dir,
env=env, env=env,
) )
timeline("sky_launch", "completed", attempt=attempt, return_code=return_code, selected_offer_id=selection.selected_offer_id)
state("REMOTE_FINISHED", selected_offer_id=selection.selected_offer_id, attempt=attempt, return_code=return_code) state("REMOTE_FINISHED", selected_offer_id=selection.selected_offer_id, attempt=attempt, return_code=return_code)
_collect_terminal_best_effort(cluster=run_id, remote_dir=config.job.artifact_dir, local_dir=local_run_dir, required=config.artifacts.required, env=env) _collect_terminal_best_effort(cluster=run_id, remote_dir=config.job.artifact_dir, local_dir=local_run_dir, required=config.artifacts.required, env=env)
status = _classify_artifacts(local_run_dir) status = _classify_artifacts(local_run_dir)
@ -447,13 +358,9 @@ def _collect_paths_with_rsync(
paths: tuple[str, ...], paths: tuple[str, ...],
env: dict[str, str], env: dict[str, str],
timeout: int, timeout: int,
required: tuple[str, ...] = (),
collection_kind: str = "artifact",
raise_on_required: bool = True,
) -> None: ) -> None:
local_dir.mkdir(parents=True, exist_ok=True) local_dir.mkdir(parents=True, exist_ok=True)
for relative_path in paths:
def copy_one(relative_path: str) -> int | None:
source = f"{cluster}:~/sky_workdir/{remote_dir}/./{relative_path}" source = f"{cluster}:~/sky_workdir/{remote_dir}/./{relative_path}"
_run_checked( _run_checked(
[ [
@ -469,49 +376,28 @@ def _collect_paths_with_rsync(
env=env, env=env,
timeout=timeout, timeout=timeout,
) )
return 0
report = collect_artifact_paths(
local_dir=local_dir,
remote_dir=f"{cluster}:~/sky_workdir/{remote_dir}",
paths=paths,
required=required,
collection_kind=collection_kind,
copy_one=copy_one,
)
failures = required_collection_failures(report, paths=paths)
if failures and raise_on_required:
names = ", ".join(str(item["expected_path"]) for item in failures)
raise RuntimeError(f"Required artifact collection failed: {names}")
def _collect_required_artifacts(*, cluster: str, remote_dir: Path, local_dir: Path, required: tuple[str, ...], env: dict[str, str]) -> None: def _collect_required_artifacts(*, cluster: str, remote_dir: Path, local_dir: Path, required: tuple[str, ...], env: dict[str, str]) -> None:
large = _large_artifact_names(required)
_collect_paths_with_rsync( _collect_paths_with_rsync(
cluster=cluster, cluster=cluster,
remote_dir=remote_dir, remote_dir=remote_dir,
local_dir=local_dir, local_dir=local_dir,
paths=large, paths=_large_artifact_names(required),
env=env, env=env,
timeout=3600, timeout=3600,
required=large,
collection_kind="large",
) )
def _collect_terminal_best_effort(*, cluster: str, remote_dir: Path, local_dir: Path, required: tuple[str, ...], env: dict[str, str]) -> None: def _collect_terminal_best_effort(*, cluster: str, remote_dir: Path, local_dir: Path, required: tuple[str, ...], env: dict[str, str]) -> None:
try: try:
terminal = _terminal_artifact_names(required)
_collect_paths_with_rsync( _collect_paths_with_rsync(
cluster=cluster, cluster=cluster,
remote_dir=remote_dir, remote_dir=remote_dir,
local_dir=local_dir, local_dir=local_dir,
paths=terminal, paths=_terminal_artifact_names(required),
env=env, env=env,
timeout=120, timeout=120,
required=tuple(name for name in terminal if name in required),
collection_kind="terminal",
raise_on_required=False,
) )
except Exception: except Exception:
_cleanup_partial_artifacts(local_dir) _cleanup_partial_artifacts(local_dir)
@ -526,8 +412,6 @@ def _collect_restart_best_effort(*, cluster: str, remote_dir: Path, local_dir: P
paths=("checkpoint_latest.pt",), paths=("checkpoint_latest.pt",),
env=env, env=env,
timeout=3600, timeout=3600,
collection_kind="restart",
raise_on_required=False,
) )
except Exception: except Exception:
_cleanup_partial_artifacts(local_dir) _cleanup_partial_artifacts(local_dir)
@ -536,14 +420,10 @@ def _collect_restart_best_effort(*, cluster: str, remote_dir: Path, local_dir: P
_LARGE_ARTIFACT_SUFFIXES = (".pt", ".pth", ".ckpt", ".safetensors") _LARGE_ARTIFACT_SUFFIXES = (".pt", ".pth", ".ckpt", ".safetensors")
_TERMINAL_ARTIFACT_NAMES = ( _TERMINAL_ARTIFACT_NAMES = (
"artifact_manifest.json", "artifact_manifest.json",
ARTIFACT_COLLECTION_REPORT,
"checksums.txt", "checksums.txt",
"config.toml", "config.toml",
"calibration_manifest.json",
"data_manifest.json", "data_manifest.json",
"disk_telemetry.json",
"environment_manifest.json", "environment_manifest.json",
"evaluation_protocol.json",
"failure_report.json", "failure_report.json",
"final_metrics.json", "final_metrics.json",
"heartbeat.json", "heartbeat.json",
@ -553,9 +433,7 @@ _TERMINAL_ARTIFACT_NAMES = (
"normalization.json", "normalization.json",
"run_manifest.json", "run_manifest.json",
"split_manifest.json", "split_manifest.json",
"startup_timeline.jsonl",
"wandb_smoke_manifest.json", "wandb_smoke_manifest.json",
"verification_report.json",
) )
@ -634,9 +512,6 @@ def _run_best_effort(argv: list[str], *, env: dict[str, str]) -> None:
def _subprocess_env() -> dict[str, str]: def _subprocess_env() -> dict[str, str]:
env = dict(os.environ) env = dict(os.environ)
env.pop("PYTHONPATH", None) env.pop("PYTHONPATH", None)
executable_dir = str(Path(sys.executable).parent)
path = env.get("PATH")
env["PATH"] = executable_dir if not path else f"{executable_dir}{os.pathsep}{path}"
return env return env
def _ensure_hf_secret_env(env: dict[str, str]) -> None: def _ensure_hf_secret_env(env: dict[str, str]) -> None:
@ -670,33 +545,6 @@ def _load_secret_env(env: dict[str, str], name: str, *, required: bool, purpose:
def _load_json_file(path: Path) -> Any:
return json.loads(path.read_text())
def _load_sky_status() -> Any:
process = subprocess.run(
["sky", "status", "--format", "json"],
check=True,
capture_output=True,
text=True,
env=_subprocess_env(),
timeout=120,
)
return json.loads(process.stdout)
def _instances_from_json_payload(payload: Any) -> list[dict[str, Any]]:
if isinstance(payload, list):
return [dict(item) for item in payload if isinstance(item, dict)]
if isinstance(payload, dict):
for key in ("instances", "results", "items"):
value = payload.get(key)
if isinstance(value, list):
return [dict(item) for item in value if isinstance(item, dict)]
raise ValueError("Vast instances JSON must be a list or contain instances/results/items")
def _emit_json(data: dict[str, Any], out: str | None) -> None: def _emit_json(data: dict[str, Any], out: str | None) -> None:
text = json.dumps(data, indent=2, sort_keys=True) + "\n" text = json.dumps(data, indent=2, sort_keys=True) + "\n"
if out: if out:
@ -705,17 +553,14 @@ def _emit_json(data: dict[str, Any], out: str | None) -> None:
print(text, end="") print(text, end="")
def _selection_from_manifest(path: Path, *, max_age_seconds: float = DEFAULT_SELECTION_MAX_AGE_SECONDS, allow_stale: bool = False) -> SelectionResult: def _selection_from_manifest(path: Path) -> SelectionResult:
from airfrans_frontier.remote.vast import VastOffer from airfrans_frontier.remote.vast import VastOffer, effective_price
data = load_selection_manifest(path, max_age_seconds=max_age_seconds, allow_stale=allow_stale) data = json.loads(path.read_text())
raw_offer = data.get("selected_offer") raw_offer = data.get("selected_offer")
if not isinstance(raw_offer, dict): if not isinstance(raw_offer, dict):
raise ValueError(f"Selection manifest missing selected_offer object: {path}") raise ValueError(f"Selection manifest missing selected_offer object: {path}")
offer = VastOffer.from_mapping({**raw_offer, "id": data.get("selected_offer_id", raw_offer.get("id"))}) offer = VastOffer.from_mapping({**raw_offer, "id": data.get("selected_offer_id", raw_offer.get("id"))})
created_at = data.get("created_at")
if not isinstance(created_at, (int, float)):
created_at = time.time()
# Preserve manifest values by building a minimal SelectionResult. Effective price is already stored. # Preserve manifest values by building a minimal SelectionResult. Effective price is already stored.
return SelectionResult( return SelectionResult(
selected_offer=offer, selected_offer=offer,
@ -724,7 +569,6 @@ def _selection_from_manifest(path: Path, *, max_age_seconds: float = DEFAULT_SEL
effective_price=float(raw_offer.get("effective_price", data.get("effective_price", 0.0))) if raw_offer else 0.0, effective_price=float(raw_offer.get("effective_price", data.get("effective_price", 0.0))) if raw_offer else 0.0,
query=data.get("query", {}) if isinstance(data.get("query"), dict) else {}, query=data.get("query", {}) if isinstance(data.get("query"), dict) else {},
policy=data.get("policy", {}) if isinstance(data.get("policy"), dict) else {}, policy=data.get("policy", {}) if isinstance(data.get("policy"), dict) else {},
created_at=float(created_at),
) )

View file

@ -1,180 +0,0 @@
from __future__ import annotations
from collections.abc import Callable, Iterable, Mapping
import json
from pathlib import Path
import time
from typing import Any
ARTIFACT_COLLECTION_REPORT = "artifact_collection_report.json"
_PARTIAL_SUFFIXES = (".tmp", ".part", ".partial")
_RSYNC_TEMP_DIRS = (".rsync-partial", ".~tmp~")
def collect_artifact_paths(
*,
local_dir: str | Path,
remote_dir: str | Path,
paths: Iterable[str],
required: Iterable[str] = (),
collection_kind: str,
copy_one: Callable[[str], int | None],
clock: Callable[[], float] | None = None,
) -> dict[str, Any]:
"""Copy artifact paths and update artifact_collection_report.json.
copy_one receives each relative artifact path. It may raise or return a
non-zero return code; both are recorded per path without losing the rest of
the collection report.
"""
now = clock or time.time
root = Path(local_dir)
root.mkdir(parents=True, exist_ok=True)
report_path = root / ARTIFACT_COLLECTION_REPORT
report = _load_report(report_path)
required_set = set(required)
attempted_paths = tuple(dict.fromkeys(paths))
batch_started_at = now()
batch_id = f"{collection_kind}-{int(batch_started_at * 1000)}-{len(report['attempts'])}"
report["batches"].append(
{
"batch_id": batch_id,
"collection_kind": collection_kind,
"started_at": batch_started_at,
"paths": list(attempted_paths),
}
)
_write_report(report_path, _refresh_summary(report, now=now()))
for relative_path in attempted_paths:
_validate_relative_path(relative_path)
source = f"{str(remote_dir).rstrip('/')}/{relative_path}"
destination = root / relative_path
started = now()
attempt: dict[str, Any] = {
"batch_id": batch_id,
"collection_kind": collection_kind,
"expected_path": relative_path,
"required": relative_path in required_set,
"source_path": source,
"local_destination": str(destination),
"attempted": True,
"started_at": started,
"bytes_copied": None,
"duration_seconds": None,
"return_code": None,
"exception": None,
"final_status": "failed",
"likely_reason": None,
}
try:
return_code = copy_one(relative_path)
if return_code is not None:
attempt["return_code"] = int(return_code)
except Exception as exc:
attempt["exception"] = {"type": type(exc).__name__, "message": str(exc)}
attempt["final_status"] = "failed"
attempt["likely_reason"] = "collection_command_failed"
else:
if attempt["return_code"] not in (None, 0):
attempt["final_status"] = "failed"
attempt["likely_reason"] = "collection_command_failed"
else:
partial = _partial_related_path(root, relative_path)
if partial is not None:
attempt["final_status"] = "partial"
attempt["likely_reason"] = "partial_or_temp_file_present"
attempt["partial_path"] = str(partial)
elif destination.is_file() and not _is_partial_name(destination.name):
attempt["final_status"] = "success"
attempt["bytes_copied"] = destination.stat().st_size
attempt["likely_reason"] = "artifact_collected"
else:
attempt["final_status"] = "missing"
attempt["likely_reason"] = "remote_missing_or_not_produced"
attempt["duration_seconds"] = max(0.0, now() - started)
report["attempts"].append(attempt)
_write_report(report_path, _refresh_summary(report, now=now()))
report["batches"][-1]["finished_at"] = now()
_write_report(report_path, _refresh_summary(report, now=now()))
return report
def required_collection_failures(report: Mapping[str, Any], *, paths: Iterable[str] | None = None) -> list[dict[str, Any]]:
selected = set(paths) if paths is not None else None
failures: list[dict[str, Any]] = []
for raw_attempt in report.get("attempts", []):
if not isinstance(raw_attempt, dict):
continue
if selected is not None and raw_attempt.get("expected_path") not in selected:
continue
if raw_attempt.get("required") and raw_attempt.get("final_status") != "success":
failures.append(dict(raw_attempt))
return failures
def _load_report(path: Path) -> dict[str, Any]:
if path.is_file():
try:
data = json.loads(path.read_text())
except json.JSONDecodeError:
data = None
if isinstance(data, dict):
data.setdefault("schema_version", 1)
data.setdefault("attempts", [])
data.setdefault("batches", [])
data.setdefault("summary", {})
return data
return {"schema_version": 1, "attempts": [], "batches": [], "summary": {}}
def _refresh_summary(report: dict[str, Any], *, now: float) -> dict[str, Any]:
counts: dict[str, int] = {}
required_missing: list[str] = []
for raw_attempt in report.get("attempts", []):
if not isinstance(raw_attempt, dict):
continue
status = str(raw_attempt.get("final_status", "unknown"))
counts[status] = counts.get(status, 0) + 1
if raw_attempt.get("required") and status != "success":
expected = raw_attempt.get("expected_path")
if isinstance(expected, str):
required_missing.append(expected)
report["summary"] = {
"updated_at": now,
"attempt_count": sum(counts.values()),
"status_counts": dict(sorted(counts.items())),
"required_uncollected": required_missing,
"ok": not required_missing,
}
return report
def _write_report(path: Path, report: Mapping[str, Any]) -> None:
path.write_text(json.dumps(report, indent=2, sort_keys=True) + "\n")
def _validate_relative_path(relative_path: str) -> None:
path = Path(relative_path)
if path.is_absolute() or ".." in path.parts:
raise ValueError(f"Artifact path must be relative and stay under artifact root: {relative_path}")
def _partial_related_path(local_dir: Path, relative_path: str) -> Path | None:
destination = local_dir / relative_path
if destination.exists() and _is_partial_name(destination.name):
return destination
for suffix in _PARTIAL_SUFFIXES:
candidate = destination.with_name(f"{destination.name}{suffix}")
if candidate.exists():
return candidate
for temp_dir in _RSYNC_TEMP_DIRS:
candidate = local_dir / temp_dir / relative_path
if candidate.exists():
return candidate
return None
def _is_partial_name(name: str) -> bool:
return name.endswith(_PARTIAL_SUFFIXES) or name in _RSYNC_TEMP_DIRS

View file

@ -1,344 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass, field
import json
from pathlib import Path
import time
from typing import Any, Callable, Iterable
QUEUED_PHASE = "queued"
HEALTHY_PHASE = "training_healthy"
COMPLETED_PHASE = "completed"
FAILED_PHASE = "failed"
FRAGILE_PHASES = frozenset(
{
"offer_selection",
"provisioning",
"cluster_startup",
"ssh_reachability",
"workdir_sync",
"environment_setup",
"data_validation",
}
)
TERMINAL_PHASES = frozenset({COMPLETED_PHASE, FAILED_PHASE})
OBSERVABILITY_EVENTS = (
"run_queued",
"capacity_acquired",
"capacity_blocked",
"offer_selected",
"provisioning_started",
"cluster_reachable",
"rsync_started",
"rsync_completed",
"setup_started",
"setup_completed",
"data_validation_started",
"data_validation_completed",
"training_healthy",
"run_completed",
"run_failed",
"cleanup_started",
"cleanup_completed",
"retry_scheduled",
"retry_exhausted",
)
_PHASE_EVENTS = {
QUEUED_PHASE: "run_queued",
"offer_selection": "capacity_acquired",
"provisioning": "provisioning_started",
"cluster_startup": "provisioning_started",
"ssh_reachability": "cluster_reachable",
"workdir_sync": "rsync_started",
"environment_setup": "setup_started",
"data_validation": "data_validation_started",
HEALTHY_PHASE: "training_healthy",
COMPLETED_PHASE: "run_completed",
FAILED_PHASE: "run_failed",
}
@dataclass(frozen=True)
class LaunchRunSpec:
run_id: str
config_path: str
@dataclass
class LaunchRunState:
run_id: str
config_path: str
phase: str = QUEUED_PHASE
selected_offer_id: int | None = None
selected_host_id: int | None = None
retry_count: int = 0
last_error: str | None = None
cleanup_state: str = "not_started"
blocked_reason: str | None = None
timestamps: dict[str, float] = field(default_factory=dict)
def to_payload(self) -> dict[str, Any]:
return {
"run_id": self.run_id,
"config_path": self.config_path,
"phase": self.phase,
"selected_offer_id": self.selected_offer_id,
"selected_host_id": self.selected_host_id,
"retry_count": self.retry_count,
"last_error": self.last_error,
"cleanup_state": self.cleanup_state,
"blocked_reason": self.blocked_reason,
"timestamps": dict(sorted(self.timestamps.items())),
}
class LaunchGroupScheduler:
"""Local launch-group state machine for bounded fragile-phase scheduling.
The scheduler does not provision machines. Callers drive phase transitions from
observed launch/training evidence and get a durable state artifact after each
transition.
"""
def __init__(
self,
run_configs: Iterable[str | Path | LaunchRunSpec],
*,
max_active: int,
max_fragile: int,
state_path: str | Path | None = None,
group_id: str | None = None,
allow_duplicate_hosts: bool = False,
clock: Callable[[], float] | None = None,
) -> None:
if max_active < 1:
raise ValueError("max_active must be >= 1")
if max_fragile < 1:
raise ValueError("max_fragile must be >= 1")
if max_fragile > max_active:
raise ValueError("max_fragile must be <= max_active")
self.clock = clock or time.time
self.group_id = group_id or f"launch-{int(self.clock())}"
self.max_active = int(max_active)
self.max_fragile = int(max_fragile)
self.allow_duplicate_hosts = bool(allow_duplicate_hosts)
self.state_path = Path(state_path) if state_path is not None else None
self.runs: dict[str, LaunchRunState] = {}
self.events: list[dict[str, Any]] = []
for spec in _coerce_run_specs(run_configs):
now = self.clock()
run = LaunchRunState(run_id=spec.run_id, config_path=spec.config_path)
run.timestamps["queued_at"] = now
self.runs[run.run_id] = run
self._record_event(run.run_id, "run_queued", phase=QUEUED_PHASE, ts=now)
if not self.runs:
raise ValueError("launch group requires at least one run config")
self.write_state()
def try_start(self, run_id: str, *, selected_offer_id: int | None = None, selected_host_id: int | None = None) -> bool:
run = self._run(run_id)
if run.phase != QUEUED_PHASE:
raise ValueError(f"Run {run_id} is not queued: {run.phase}")
blocker = self._capacity_blocker(selected_host_id=selected_host_id)
if blocker is not None:
run.blocked_reason = blocker
self._record_event(run_id, "capacity_blocked", phase=run.phase, reason=blocker)
self.write_state()
return False
run.phase = "offer_selection"
run.blocked_reason = None
run.selected_offer_id = selected_offer_id
run.selected_host_id = selected_host_id
now = self.clock()
run.timestamps["capacity_acquired_at"] = now
run.timestamps["offer_selection_at"] = now
self._record_event(run_id, "capacity_acquired", phase=run.phase, ts=now)
if selected_offer_id is not None or selected_host_id is not None:
self._record_event(
run_id,
"offer_selected",
phase=run.phase,
selected_offer_id=selected_offer_id,
selected_host_id=selected_host_id,
)
self.write_state()
return True
def assign_offer(self, run_id: str, *, selected_offer_id: int, selected_host_id: int | None) -> bool:
run = self._run(run_id)
if run.phase == QUEUED_PHASE:
return self.try_start(run_id, selected_offer_id=selected_offer_id, selected_host_id=selected_host_id)
if run.phase in TERMINAL_PHASES:
raise ValueError(f"Cannot assign offer to terminal run {run_id}: {run.phase}")
if self._host_collision(selected_host_id, excluding_run_id=run_id):
run.blocked_reason = "host_collision"
self._record_event(run_id, "capacity_blocked", phase=run.phase, reason="host_collision", selected_host_id=selected_host_id)
self.write_state()
return False
run.selected_offer_id = int(selected_offer_id)
run.selected_host_id = selected_host_id
run.blocked_reason = None
run.timestamps["offer_selected_at"] = self.clock()
self._record_event(
run_id,
"offer_selected",
phase=run.phase,
selected_offer_id=selected_offer_id,
selected_host_id=selected_host_id,
)
self.write_state()
return True
def transition(self, run_id: str, phase: str, *, error: str | None = None, cleanup_state: str | None = None) -> None:
run = self._run(run_id)
run.phase = phase
run.blocked_reason = None
if error is not None:
run.last_error = error
if cleanup_state is not None:
run.cleanup_state = cleanup_state
now = self.clock()
run.timestamps[f"{phase}_at"] = now
self._record_event(run_id, _PHASE_EVENTS.get(phase, phase), phase=phase, ts=now, error=error, cleanup_state=cleanup_state)
self.write_state()
def mark_training_healthy(self, run_id: str) -> None:
self.transition(run_id, HEALTHY_PHASE)
def complete_run(self, run_id: str) -> None:
self.transition(run_id, COMPLETED_PHASE)
def fail_run(self, run_id: str, error: str) -> None:
self.transition(run_id, FAILED_PHASE, error=error)
def schedule_retry(self, run_id: str, error: str) -> None:
run = self._run(run_id)
run.retry_count += 1
run.last_error = error
run.phase = QUEUED_PHASE
run.blocked_reason = None
run.timestamps["retry_scheduled_at"] = self.clock()
self._record_event(run_id, "retry_scheduled", phase=run.phase, retry_count=run.retry_count, error=error)
self.write_state()
def capacity_snapshot(self) -> dict[str, int]:
return {
"active": self._active_count(),
"fragile": self._fragile_count(),
"pending": len(self._runs_in_phase(QUEUED_PHASE)),
"healthy": len(self._runs_in_phase(HEALTHY_PHASE)),
"completed": len(self._runs_in_phase(COMPLETED_PHASE)),
"failed": len(self._runs_in_phase(FAILED_PHASE)),
}
def to_payload(self) -> dict[str, Any]:
pending = self._runs_in_phase(QUEUED_PHASE)
healthy = self._runs_in_phase(HEALTHY_PHASE)
completed = self._runs_in_phase(COMPLETED_PHASE)
failed = self._runs_in_phase(FAILED_PHASE)
running = [
run_id
for run_id, run in self.runs.items()
if run.phase not in {QUEUED_PHASE, HEALTHY_PHASE, COMPLETED_PHASE, FAILED_PHASE}
]
return {
"schema_version": 1,
"launch_group_id": self.group_id,
"requested_run_configs": [run.config_path for run in self.runs.values()],
"limits": {
"max_active": self.max_active,
"max_fragile": self.max_fragile,
"allow_duplicate_hosts": self.allow_duplicate_hosts,
"fragile_phases": sorted(FRAGILE_PHASES),
},
"pending_runs": pending,
"running_runs": running,
"healthy_runs": healthy,
"completed_runs": completed,
"failed_runs": failed,
"counts": self.capacity_snapshot(),
"phase_counts": self._phase_counts(),
"runs": {run_id: run.to_payload() for run_id, run in self.runs.items()},
"events": list(self.events),
"updated_at": self.clock(),
}
def write_state(self) -> Path | None:
if self.state_path is None:
return None
self.state_path.parent.mkdir(parents=True, exist_ok=True)
self.state_path.write_text(json.dumps(self.to_payload(), indent=2, sort_keys=True) + "\n")
return self.state_path
def _capacity_blocker(self, *, selected_host_id: int | None) -> str | None:
if self._active_count() >= self.max_active:
return "max_active"
if self._fragile_count() >= self.max_fragile:
return "max_fragile"
if self._host_collision(selected_host_id):
return "host_collision"
return None
def _host_collision(self, selected_host_id: int | None, *, excluding_run_id: str | None = None) -> bool:
if selected_host_id is None or self.allow_duplicate_hosts:
return False
for run_id, run in self.runs.items():
if run_id == excluding_run_id:
continue
if run.phase not in FRAGILE_PHASES:
continue
if run.selected_host_id == selected_host_id:
return True
return False
def _active_count(self) -> int:
return sum(1 for run in self.runs.values() if run.phase not in {QUEUED_PHASE, *TERMINAL_PHASES})
def _fragile_count(self) -> int:
return sum(1 for run in self.runs.values() if run.phase in FRAGILE_PHASES)
def _runs_in_phase(self, phase: str) -> list[str]:
return [run_id for run_id, run in self.runs.items() if run.phase == phase]
def _phase_counts(self) -> dict[str, int]:
counts: dict[str, int] = {}
for run in self.runs.values():
counts[run.phase] = counts.get(run.phase, 0) + 1
return dict(sorted(counts.items()))
def _run(self, run_id: str) -> LaunchRunState:
try:
return self.runs[run_id]
except KeyError as exc:
raise KeyError(f"Unknown launch run id: {run_id}") from exc
def _record_event(self, run_id: str, event: str, *, phase: str, ts: float | None = None, **fields: Any) -> None:
record = {
"launch_group_id": self.group_id,
"run_id": run_id,
"event": event,
"phase": phase,
"ts": self.clock() if ts is None else ts,
**{key: value for key, value in fields.items() if value is not None},
}
self.events.append(record)
def _coerce_run_specs(run_configs: Iterable[str | Path | LaunchRunSpec]) -> list[LaunchRunSpec]:
result: list[LaunchRunSpec] = []
seen: set[str] = set()
for index, item in enumerate(run_configs, start=1):
if isinstance(item, LaunchRunSpec):
spec = item
else:
config_path = str(Path(item))
base = Path(config_path).stem.replace("_", "-") or f"run-{index}"
run_id = base if base not in seen else f"{base}-{index}"
spec = LaunchRunSpec(run_id=run_id, config_path=config_path)
if spec.run_id in seen:
raise ValueError(f"Duplicate run id in launch group: {spec.run_id}")
seen.add(spec.run_id)
result.append(spec)
return result

View file

@ -1,120 +0,0 @@
from __future__ import annotations
from datetime import UTC, datetime
import json
from pathlib import Path
import time
from typing import Any, Mapping
DEFAULT_SELECTION_MAX_AGE_SECONDS = 15 * 60
def selection_freshness_report(
manifest: Mapping[str, Any],
*,
max_age_seconds: float = DEFAULT_SELECTION_MAX_AGE_SECONDS,
now: float | None = None,
) -> dict[str, Any]:
"""Return freshness metadata for a Vast offer selection artifact."""
if max_age_seconds < 0:
raise ValueError("max_age_seconds must be non-negative")
checked_at = time.time() if now is None else float(now)
created_at = _created_at_seconds(manifest)
if created_at is None:
return {
"created_at": None,
"created_at_iso": None,
"checked_at": checked_at,
"age_seconds": None,
"max_age_seconds": float(max_age_seconds),
"is_fresh": False,
"reason": "missing_creation_time",
}
age = max(0.0, checked_at - created_at)
is_fresh = age <= max_age_seconds
return {
"created_at": created_at,
"created_at_iso": datetime.fromtimestamp(created_at, UTC).isoformat(),
"checked_at": checked_at,
"age_seconds": age,
"max_age_seconds": float(max_age_seconds),
"is_fresh": is_fresh,
"reason": "fresh" if is_fresh else "stale",
}
def require_fresh_selection(
manifest: Mapping[str, Any],
*,
max_age_seconds: float = DEFAULT_SELECTION_MAX_AGE_SECONDS,
now: float | None = None,
path: str | Path | None = None,
) -> dict[str, Any]:
report = selection_freshness_report(manifest, max_age_seconds=max_age_seconds, now=now)
if not report["is_fresh"]:
location = f" {path}" if path is not None else ""
age = report["age_seconds"]
if age is None:
raise ValueError(f"Selection artifact{location} has no creation time and is stale by policy")
raise ValueError(
f"Selection artifact{location} is stale: age_seconds={age:.3f} "
f"max_age_seconds={float(max_age_seconds):.3f}"
)
return report
def load_selection_manifest(
path: str | Path,
*,
max_age_seconds: float = DEFAULT_SELECTION_MAX_AGE_SECONDS,
allow_stale: bool = False,
now: float | None = None,
) -> dict[str, Any]:
manifest_path = Path(path)
data = json.loads(manifest_path.read_text())
if not isinstance(data, dict):
raise ValueError(f"Selection manifest is not a JSON object: {manifest_path}")
report = selection_freshness_report(data, max_age_seconds=max_age_seconds, now=now)
data["freshness"] = report
if not allow_stale:
require_fresh_selection(data, max_age_seconds=max_age_seconds, now=now, path=manifest_path)
return data
def _created_at_seconds(manifest: Mapping[str, Any]) -> float | None:
for key in ("created_at", "selected_at", "creation_time"):
value = manifest.get(key)
parsed = _parse_timestamp_seconds(value)
if parsed is not None:
return parsed
for key in ("created_at_iso", "selected_at_iso", "creation_time_iso"):
value = manifest.get(key)
parsed = _parse_timestamp_seconds(value)
if parsed is not None:
return parsed
return None
def _parse_timestamp_seconds(value: object) -> float | None:
if isinstance(value, bool) or value is None:
return None
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
raw = value.strip()
if not raw:
return None
try:
return float(raw)
except ValueError:
pass
try:
normalized = raw[:-1] + "+00:00" if raw.endswith("Z") else raw
parsed = datetime.fromisoformat(normalized)
except ValueError:
return None
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=UTC)
return parsed.timestamp()
return None

View file

@ -32,7 +32,6 @@ def render_skypilot_yaml(
env_lines = [ env_lines = [
"envs:", "envs:",
f" AIRFRANS_REMOTE_RUN_ID: {run_id}", f" AIRFRANS_REMOTE_RUN_ID: {run_id}",
f" AIRFRANS_STARTUP_TIMELINE: {_yaml_scalar(str(config.job.artifact_dir / 'startup_timeline.jsonl'))}",
] ]
if resume_checkpoint is not None: if resume_checkpoint is not None:
env_lines.append(f" AIRFRANS_RESUME_CHECKPOINT: {_yaml_scalar(str(resume_checkpoint))}") env_lines.append(f" AIRFRANS_RESUME_CHECKPOINT: {_yaml_scalar(str(resume_checkpoint))}")
@ -74,66 +73,9 @@ def _compose_setup(config: RemoteRunConfig) -> str:
[ [
"set -euo pipefail", "set -euo pipefail",
"export PATH=\"$HOME/.local/bin:$PATH\"", "export PATH=\"$HOME/.local/bin:$PATH\"",
f"mkdir -p {_sh_quote(str(config.job.artifact_dir))}",
_timeline_shell_function(),
_timeline_event("setup", "started"),
"if ! command -v uv >/dev/null 2>&1; then curl -LsSf https://astral.sh/uv/install.sh | sh; fi", "if ! command -v uv >/dev/null 2>&1; then curl -LsSf https://astral.sh/uv/install.sh | sh; fi",
"export PATH=\"$HOME/.local/bin:$PATH\"", "export PATH=\"$HOME/.local/bin:$PATH\"",
_timeline_event("disk_preflight", "started"),
_remote_disk_preflight(config),
_timeline_event("disk_preflight", "completed"),
_timeline_event("bootstrap", "started"),
config.bootstrap.command.strip(), config.bootstrap.command.strip(),
_timeline_event("bootstrap", "completed"),
_timeline_event("setup", "completed"),
]
)
def _remote_disk_preflight(config: RemoteRunConfig) -> str:
requested_gb = config.provider.disk_gb
minimum_total_kib = int(requested_gb * 1024 * 1024 * 0.90)
telemetry_path = config.job.artifact_dir / "disk_telemetry.json"
return "\n".join(
[
"echo 'airfrans_disk_df_start'",
"df -h .",
f"airfrans_disk_requested_gb={requested_gb}",
"airfrans_disk_total_kib=$(df -Pk . | tail -n 1 | tr -s ' ' | cut -d ' ' -f 2)",
"airfrans_disk_available_kib=$(df -Pk . | tail -n 1 | tr -s ' ' | cut -d ' ' -f 4)",
f"airfrans_disk_minimum_requested_total_kib={minimum_total_kib}",
"echo \"airfrans_disk_requested_gb=${airfrans_disk_requested_gb}\"",
"echo \"airfrans_disk_total_kib=${airfrans_disk_total_kib}\"",
"echo \"airfrans_disk_available_kib=${airfrans_disk_available_kib}\"",
"echo \"airfrans_disk_minimum_requested_total_kib=${airfrans_disk_minimum_requested_total_kib}\"",
"airfrans_disk_capacity_policy=backpressure_adaptive",
f"if [ \"$airfrans_disk_total_kib\" -lt {minimum_total_kib} ]; then",
f" echo \"warning: effective filesystem total ${{airfrans_disk_total_kib}} KiB is below 90% of requested {requested_gb}GB disk; continuing because runtime cache backpressure can adapt\" >&2",
" airfrans_disk_capacity_status=below_requested",
"else",
" airfrans_disk_capacity_status=ok",
"fi",
f"python3 - \"$airfrans_disk_requested_gb\" \"$airfrans_disk_total_kib\" \"$airfrans_disk_available_kib\" \"$airfrans_disk_minimum_requested_total_kib\" \"$airfrans_disk_capacity_status\" {_sh_quote(str(telemetry_path))} <<'PY'",
"import json, os, sys, time",
"requested_gb, total_kib, available_kib, minimum_total_kib, status, path = sys.argv[1:7]",
"payload = {",
" 'schema_version': 1,",
" 'recorded_at': time.time(),",
" 'requested_gb': int(requested_gb),",
" 'total_kib': int(total_kib),",
" 'available_kib': int(available_kib),",
" 'minimum_requested_total_kib': int(minimum_total_kib),",
" 'capacity_status': status,",
" 'capacity_policy': 'backpressure_adaptive',",
" 'hard_failed': False,",
"}",
"directory = os.path.dirname(path)",
"if directory:",
" os.makedirs(directory, exist_ok=True)",
"with open(path, 'w', encoding='utf-8') as handle:",
" json.dump(payload, handle, indent=2, sort_keys=True)",
" handle.write('\\n')",
"PY",
] ]
) )
@ -143,63 +85,16 @@ def _compose_run(config: RemoteRunConfig, *, run_id: str) -> str:
"set -euo pipefail", "set -euo pipefail",
"export PATH=\"$HOME/.local/bin:$PATH\"", "export PATH=\"$HOME/.local/bin:$PATH\"",
f"mkdir -p {_sh_quote(str(config.job.artifact_dir))}", f"mkdir -p {_sh_quote(str(config.job.artifact_dir))}",
_timeline_shell_function(),
_timeline_event("run", "started"),
_timeline_event("gpu_probe", "started"),
f"nvidia-smi | tee {_sh_quote(str(config.job.artifact_dir / 'nvidia_smi.txt'))}", f"nvidia-smi | tee {_sh_quote(str(config.job.artifact_dir / 'nvidia_smi.txt'))}",
_timeline_event("gpu_probe", "completed"),
] ]
if config.data.validation_command: if config.data.validation_command:
lines.extend( lines.append(config.data.validation_command.strip())
[ lines.append(config.job.command.strip())
_timeline_event("data_validation", "started"), lines.append(f"uv run --no-dev remote-run verify-artifacts {_sh_quote(str(config.job.artifact_dir))}")
config.data.validation_command.strip(), lines.append(f"echo 'remote run {run_id} complete'")
_timeline_event("data_validation", "completed"),
]
)
lines.extend(
[
_timeline_event("training_command", "started"),
config.job.command.strip(),
_timeline_event("training_command", "completed"),
_timeline_event("artifact_verification", "started"),
f"uv run --no-dev remote-run verify-artifacts {_sh_quote(str(config.job.artifact_dir))}",
_timeline_event("artifact_verification", "completed"),
_timeline_event("run", "completed"),
f"echo 'remote run {run_id} complete'",
]
)
return "\n".join(lines) return "\n".join(lines)
def _timeline_shell_function() -> str:
return "\n".join(
[
"airfrans_timeline() {",
" python3 - \"$1\" \"$2\" <<'PY'",
"import json, os, sys, time",
"path = os.environ.get('AIRFRANS_STARTUP_TIMELINE', 'artifacts/current_run/startup_timeline.jsonl')",
"record = {",
" 'run_id': os.environ.get('AIRFRANS_REMOTE_RUN_ID'),",
" 'ts': time.time(),",
" 'phase': sys.argv[1],",
" 'event': sys.argv[2],",
"}",
"directory = os.path.dirname(path)",
"if directory:",
" os.makedirs(directory, exist_ok=True)",
"with open(path, 'a', encoding='utf-8') as handle:",
" handle.write(json.dumps(record, sort_keys=True) + '\\n')",
"PY",
"}",
]
)
def _timeline_event(phase: str, event: str) -> str:
return f"airfrans_timeline {_sh_quote(phase)} {_sh_quote(event)}"
def _accelerator(config: RemoteRunConfig) -> str: def _accelerator(config: RemoteRunConfig) -> str:
name = config.provider.gpu.name or "T4" name = config.provider.gpu.name or "T4"
aliases = {"Tesla T4": "T4", "RTX 3060 Ti": "RTX3060"} aliases = {"Tesla T4": "T4", "RTX 3060 Ti": "RTX3060"}

View file

@ -52,25 +52,8 @@ def run_smoke_training(
previous_run_id = os.environ.get("AIRFRANS_REMOTE_RUN_ID") previous_run_id = os.environ.get("AIRFRANS_REMOTE_RUN_ID")
os.environ["AIRFRANS_OBSERVABILITY_DIR"] = str(output_dir) os.environ["AIRFRANS_OBSERVABILITY_DIR"] = str(output_dir)
os.environ["AIRFRANS_REMOTE_RUN_ID"] = run_id os.environ["AIRFRANS_REMOTE_RUN_ID"] = run_id
training_dir: Path | None = None
error: Exception | None = None
try: try:
result = train_from_config_path(config_path, resume_path=resume_path or os.environ.get("AIRFRANS_RESUME_CHECKPOINT")) result = train_from_config_path(config_path, resume_path=resume_path or os.environ.get("AIRFRANS_RESUME_CHECKPOINT"))
training_dir = result.run_dir
except Exception as exc:
error = exc
training_dir = _latest_training_run_dir(config_path)
if training_dir is None:
_write_json(
output_dir / "failure_report.json",
{
"run_id": run_id,
"phase": "training",
"error_type": type(exc).__name__,
"error_message": str(exc),
"timestamp": time.time(),
},
)
finally: finally:
if previous_observability_dir is None: if previous_observability_dir is None:
os.environ.pop("AIRFRANS_OBSERVABILITY_DIR", None) os.environ.pop("AIRFRANS_OBSERVABILITY_DIR", None)
@ -82,45 +65,47 @@ def run_smoke_training(
os.environ["AIRFRANS_REMOTE_RUN_ID"] = previous_run_id os.environ["AIRFRANS_REMOTE_RUN_ID"] = previous_run_id
finished = time.time() finished = time.time()
if training_dir is not None: training_dir = result.run_dir
_copy_training_artifacts(training_dir, output_dir) required_from_training = [
if error is not None and not (output_dir / "failure_report.json").is_file(): "final_metrics.json",
_write_json( "metrics.jsonl",
output_dir / "failure_report.json", "latest_metrics.json",
{ "heartbeat.json",
"run_id": run_id, "checkpoint_latest.pt",
"phase": "training", "checkpoint_best.pt",
"error_type": type(error).__name__, "checkpoint_final.pt",
"error_message": str(error), "config.toml",
"training_run_dir": str(training_dir) if training_dir is not None else None, "normalization.json",
"timestamp": time.time(), "split_manifest.json",
}, ]
) for name in required_from_training:
source = training_dir / name
if source.is_file():
shutil.copy2(source, output_dir / name)
latest_metrics = _read_json(output_dir / "latest_metrics.json")
run_manifest: dict[str, Any] = { run_manifest: dict[str, Any] = {
"run_id": run_id, "run_id": run_id,
"command": f"remote-run smoke-train {config_path}", "command": f"remote-run smoke-train {config_path}",
"started_at": started, "started_at": started,
"finished_at": finished, "finished_at": finished,
"elapsed_seconds": finished - started, "elapsed_seconds": finished - started,
"exit_code": 0 if error is None else 1, "exit_code": 0,
"training_run_dir": str(training_dir) if training_dir is not None else None, "training_run_dir": str(training_dir),
"artifact_dir": str(output_dir), "artifact_dir": str(output_dir),
"final_metrics_path": str(output_dir / "final_metrics.json") if (output_dir / "final_metrics.json").is_file() else None, "final_metrics_path": str(output_dir / "final_metrics.json"),
"failure_report_path": str(output_dir / "failure_report.json") if (output_dir / "failure_report.json").is_file() else None, "checkpoint_path": str(output_dir / "checkpoint_latest.pt"),
"checkpoint_path": str(output_dir / "checkpoint_latest.pt") if (output_dir / "checkpoint_latest.pt").is_file() else None,
"resume_path": str(resume_path) if resume_path is not None else None, "resume_path": str(resume_path) if resume_path is not None else None,
} }
_write_json(output_dir / "run_manifest.json", run_manifest) _write_json(output_dir / "run_manifest.json", run_manifest)
latest_metrics = _read_json(output_dir / "latest_metrics.json")
_write_json( _write_json(
heartbeat_path, heartbeat_path,
{ {
"run_id": run_id, "run_id": run_id,
"phase": "completed" if error is None else "failed", "phase": "completed",
"epoch": latest_metrics.get("epoch"), "epoch": latest_metrics.get("epoch"),
"step": latest_metrics.get("step"), "step": latest_metrics.get("step"),
"latest_checkpoint": "checkpoint_final.pt" if error is None else "checkpoint_latest.pt", "latest_checkpoint": "checkpoint_final.pt",
"latest_metrics": latest_metrics, "latest_metrics": latest_metrics,
"started_at": started, "started_at": started,
"finished_at": finished, "finished_at": finished,
@ -128,89 +113,8 @@ def run_smoke_training(
"timestamp": time.time(), "timestamp": time.time(),
}, },
) )
if error is None: verify_artifacts(output_dir)
verify_artifacts(output_dir, required=_smoke_required(success=True))
return output_dir return output_dir
verify_artifacts(output_dir, required=_smoke_required(success=False))
raise error
def _copy_training_artifacts(training_dir: Path, output_dir: Path) -> None:
names = (
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"failure_report.json",
"normalization.json",
"split_manifest.json",
"data_manifest.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
"streaming_events.jsonl",
"streaming_state.json",
"streaming_summary.json",
"processed_upload_manifest.json",
)
for name in names:
source = training_dir / name
if source.is_file():
shutil.copy2(source, output_dir / name)
def _latest_training_run_dir(config_path: str | Path) -> Path | None:
try:
from airfrans_frontier.training.config import load_training_config
config = load_training_config(config_path)
except Exception:
return None
root = config.run.artifact_dir
if not root.is_dir():
return None
candidates = [path for path in root.iterdir() if path.is_dir()]
if not candidates:
return None
return max(candidates, key=lambda path: path.stat().st_mtime)
def _smoke_required(*, success: bool) -> tuple[str, ...]:
if not success:
return (
"heartbeat.json",
"environment_manifest.json",
"run_manifest.json",
"failure_report.json",
)
return (
"config.toml",
"metrics.jsonl",
"latest_metrics.json",
"heartbeat.json",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"normalization.json",
"split_manifest.json",
"data_manifest.json",
"environment_manifest.json",
"calibration_manifest.json",
"evaluation_protocol.json",
"hf_upload_manifest.json",
"run_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"checkpoint_final.pt",
"final_metrics.json",
)
def run_hf_upload_smoke( def run_hf_upload_smoke(
*, *,

View file

@ -1,15 +1,12 @@
from __future__ import annotations from __future__ import annotations
from datetime import UTC, datetime
import json import json
import math import math
import os import os
import time
import urllib.error
import urllib.parse import urllib.parse
import urllib.request import urllib.request
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass
from typing import Any, Iterable, Mapping from typing import Any, Mapping
from airfrans_frontier.remote.config import RemoteRunConfig, SelectionConfig from airfrans_frontier.remote.config import RemoteRunConfig, SelectionConfig
@ -20,7 +17,6 @@ class VastOffer:
gpu_name: str gpu_name: str
dph_total: float dph_total: float
gpu_ram: float | None gpu_ram: float | None
disk_space: float | None
geolocation: str | None geolocation: str | None
inet_down_cost_per_tb: float inet_down_cost_per_tb: float
inet_up_cost_per_tb: float inet_up_cost_per_tb: float
@ -40,7 +36,6 @@ class VastOffer:
gpu_name=_string(data, "gpu_name"), gpu_name=_string(data, "gpu_name"),
dph_total=_float(data, "dph_total"), dph_total=_float(data, "dph_total"),
gpu_ram=_optional_float(data, "gpu_ram"), gpu_ram=_optional_float(data, "gpu_ram"),
disk_space=_optional_float(data, "disk_space"),
geolocation=_optional_string(data, "geolocation"), geolocation=_optional_string(data, "geolocation"),
inet_down_cost_per_tb=_optional_float(data, "internet_down_cost_per_tb") or 0.0, inet_down_cost_per_tb=_optional_float(data, "internet_down_cost_per_tb") or 0.0,
inet_up_cost_per_tb=_optional_float(data, "internet_up_cost_per_tb") or 0.0, inet_up_cost_per_tb=_optional_float(data, "internet_up_cost_per_tb") or 0.0,
@ -63,7 +58,6 @@ class SelectionResult:
effective_price: float effective_price: float
query: dict[str, Any] query: dict[str, Any]
policy: dict[str, Any] policy: dict[str, Any]
created_at: float = field(default_factory=time.time)
@property @property
def selected_offer_id(self) -> int: def selected_offer_id(self) -> int:
@ -72,7 +66,6 @@ class SelectionResult:
def to_manifest(self) -> dict[str, Any]: def to_manifest(self) -> dict[str, Any]:
offer = asdict(self.selected_offer) offer = asdict(self.selected_offer)
offer["effective_price"] = self.effective_price offer["effective_price"] = self.effective_price
now = time.time()
return { return {
"selected_offer_id": self.selected_offer_id, "selected_offer_id": self.selected_offer_id,
"selected_offer": offer, "selected_offer": offer,
@ -80,9 +73,6 @@ class SelectionResult:
"survivor_count": self.survivor_count, "survivor_count": self.survivor_count,
"query": self.query, "query": self.query,
"policy": self.policy, "policy": self.policy,
"created_at": self.created_at,
"created_at_iso": datetime.fromtimestamp(self.created_at, UTC).isoformat(),
"age_seconds": max(0.0, now - self.created_at),
} }
@ -121,7 +111,6 @@ def build_query(config: RemoteRunConfig) -> dict[str, Any]:
query["verified"] = {"eq": True} query["verified"] = {"eq": True}
if provider.gpu.min_vram_gb is not None: if provider.gpu.min_vram_gb is not None:
query["gpu_ram"] = {"gte": provider.gpu.min_vram_gb * 1024} query["gpu_ram"] = {"gte": provider.gpu.min_vram_gb * 1024}
query["disk_space"] = {"gte": provider.disk_gb}
if provider.gpu.name: if provider.gpu.name:
query["gpu_name"] = {"eq": provider.gpu.name} query["gpu_name"] = {"eq": provider.gpu.name}
return query return query
@ -145,106 +134,22 @@ def search_offers(*, base_url: str, api_key: str, query: Mapping[str, Any]) -> l
raise RuntimeError("Vast offer search response missing offers list") raise RuntimeError("Vast offer search response missing offers list")
return [VastOffer.from_mapping(item) for item in raw_offers if isinstance(item, Mapping)] return [VastOffer.from_mapping(item) for item in raw_offers if isinstance(item, Mapping)]
def list_instances(*, base_url: str, api_key: str) -> list[dict[str, Any]]:
payload = _vast_api_json_request(
base_url=base_url,
api_key=api_key,
path="/api/v0/instances/",
method="GET",
)
return _instances_from_payload(payload)
def choose_offer(offers: list[VastOffer], config: RemoteRunConfig, *, query: Mapping[str, Any]) -> SelectionResult:
def destroy_instance(*, base_url: str, api_key: str, instance_id: int) -> Any:
return _vast_api_json_request(
base_url=base_url,
api_key=api_key,
path=f"/api/v0/instances/{int(instance_id)}/",
method="DELETE",
)
def summarize_instances(instances: list[Mapping[str, Any]]) -> list[dict[str, Any]]:
fields = (
"id",
"instance_id",
"machine_id",
"host_id",
"label",
"status",
"actual_status",
"gpu_name",
"num_gpus",
"dph_total",
"ssh_host",
"ssh_port",
"start_date",
)
summaries: list[dict[str, Any]] = []
for instance in instances:
summary = {field: instance[field] for field in fields if field in instance}
summaries.append(summary)
return summaries
def _vast_api_json_request(*, base_url: str, api_key: str, path: str, method: str) -> Any:
url = f"{base_url.rstrip('/')}/{path.lstrip('/')}"
request = urllib.request.Request(url, headers={"Authorization": f"Bearer {api_key}"}, method=method)
try:
with urllib.request.urlopen(request, timeout=45) as response:
return json.loads(response.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")
raise RuntimeError(f"Vast API {method} {path} HTTP {exc.code}: {body}") from exc
except OSError as exc:
raise RuntimeError(f"Vast API {method} {path} failed: {exc}") from exc
def _instances_from_payload(payload: Any) -> list[dict[str, Any]]:
if isinstance(payload, list):
raw_instances = payload
elif isinstance(payload, Mapping):
raw_instances = None
for key in ("instances", "results", "items"):
value = payload.get(key)
if isinstance(value, list):
raw_instances = value
break
if raw_instances is None:
raise RuntimeError("Vast instances response missing instances list")
else:
raise RuntimeError("Vast instances response is not JSON object or list")
return [dict(item) for item in raw_instances if isinstance(item, Mapping)]
def choose_offer(
offers: list[VastOffer],
config: RemoteRunConfig,
*,
query: Mapping[str, Any],
reserved_host_ids: Iterable[int] = (),
allow_reserved_hosts: bool = False,
) -> SelectionResult:
survivors = reachable_offers(offers, config.selection) survivors = reachable_offers(offers, config.selection)
ranked = rank_survivors(survivors, config.selection) ranked = rank_survivors(survivors, config.selection)
if config.provider.max_price_per_hour is not None: if config.provider.max_price_per_hour is not None:
ranked = [offer for offer in ranked if effective_price(offer, config.selection) <= config.provider.max_price_per_hour] ranked = [offer for offer in ranked if effective_price(offer, config.selection) <= config.provider.max_price_per_hour]
reserved_hosts = set(reserved_host_ids)
if reserved_hosts and not allow_reserved_hosts:
ranked = [offer for offer in ranked if offer.host_id is None or offer.host_id not in reserved_hosts]
if not ranked: if not ranked:
raise RuntimeError("No Vast offers survived quality, price, and host anti-collision filters") raise RuntimeError("No Vast offers survived quality filters and price cap")
selected = ranked[0] selected = ranked[0]
policy = selection_policy_manifest(config)
policy["reserved_host_ids"] = sorted(reserved_hosts)
policy["allow_reserved_hosts"] = bool(allow_reserved_hosts)
return SelectionResult( return SelectionResult(
selected_offer=selected, selected_offer=selected,
candidate_count=len(offers), candidate_count=len(offers),
survivor_count=len(ranked), survivor_count=len(ranked),
effective_price=effective_price(selected, config.selection), effective_price=effective_price(selected, config.selection),
query=dict(query), query=dict(query),
policy=policy, policy=selection_policy_manifest(config),
) )

View file

@ -144,7 +144,7 @@ def _artifact_manifest_text(root: Path) -> tuple[str, str]:
files = sorted( files = sorted(
path path
for path in root.rglob("*") for path in root.rglob("*")
if path.is_file() and path.name not in {"artifact_manifest.json", "checksums.txt", "verification_report.json"} if path.is_file() and path.name not in {"artifact_manifest.json", "checksums.txt"}
) )
manifest = { manifest = {
"artifact_dir": str(root), "artifact_dir": str(root),

View file

@ -1,138 +0,0 @@
from __future__ import annotations
import json
import tempfile
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from airfrans_frontier.training.config import TrainingConfig
def estimate_forward_flops_per_item(model: nn.Module) -> int:
total = 0
for module in model.modules():
if isinstance(module, nn.Linear):
total += 2 * module.in_features * module.out_features
if module.bias is not None:
total += module.out_features
return int(total)
def estimate_training_compute(*, steps: int, batch_size: int, forward_flops_per_item: int) -> int:
return int(steps * batch_size * forward_flops_per_item * 3)
def checkpoint_size_bytes(run_dir: Path, name: str) -> int | None:
path = run_dir / name
if not path.is_file():
return None
return int(path.stat().st_size)
def gpu_memory_metrics(device: torch.device) -> dict[str, int | None]:
if device.type != "cuda":
return {
"gpu_memory_allocated_mb": None,
"gpu_memory_reserved_mb": None,
"gpu_memory_peak_allocated_mb": None,
}
index = device.index if device.index is not None else torch.cuda.current_device()
return {
"gpu_memory_allocated_mb": int(torch.cuda.memory_allocated(index) // (1024 * 1024)),
"gpu_memory_reserved_mb": int(torch.cuda.memory_reserved(index) // (1024 * 1024)),
"gpu_memory_peak_allocated_mb": int(torch.cuda.max_memory_allocated(index) // (1024 * 1024)),
}
def measure_training_step(
model: nn.Module,
optimizer: torch.optim.Optimizer,
features: np.ndarray,
targets: np.ndarray,
*,
batch_size: int,
steps: int,
device: torch.device,
) -> dict[str, Any]:
if steps <= 0:
raise ValueError("Calibration steps must be positive")
rng = np.random.default_rng(1729)
model.train()
if device.type == "cuda":
torch.cuda.reset_peak_memory_stats(device)
torch.cuda.synchronize(device)
started = time.perf_counter()
last_loss = 0.0
for _ in range(steps):
indices = rng.integers(0, features.shape[0], size=batch_size)
batch_features = torch.from_numpy(np.ascontiguousarray(features[indices], dtype=np.float32)).to(device)
batch_targets = torch.from_numpy(np.ascontiguousarray(targets[indices], dtype=np.float32)).to(device)
optimizer.zero_grad(set_to_none=True)
predictions = model(batch_features)
loss = F.mse_loss(predictions, batch_targets)
loss.backward()
optimizer.step()
last_loss = float(loss.detach().cpu().item())
if device.type == "cuda":
torch.cuda.synchronize(device)
elapsed = time.perf_counter() - started
return {
"calibration_steps": steps,
"step_time_seconds": elapsed / steps,
"points_per_sec": steps * batch_size / max(elapsed, 1e-12),
"last_calibration_loss": last_loss,
**gpu_memory_metrics(device),
}
def measure_validation_runtime(
model: nn.Module,
features: np.ndarray,
*,
batch_size: int,
device: torch.device,
) -> float:
model.eval()
if device.type == "cuda":
torch.cuda.synchronize(device)
started = time.perf_counter()
with torch.no_grad():
for start in range(0, features.shape[0], batch_size):
stop = min(start + batch_size, features.shape[0])
batch_features = torch.from_numpy(np.ascontiguousarray(features[start:stop], dtype=np.float32)).to(device)
model(batch_features)
if device.type == "cuda":
torch.cuda.synchronize(device)
return time.perf_counter() - started
def measure_checkpoint_size(payload: dict[str, Any]) -> int:
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "checkpoint.pt"
torch.save(payload, path)
return int(path.stat().st_size)
def write_calibration_report(path: str | Path, data: dict[str, Any]) -> Path:
report_path = Path(path)
report_path.parent.mkdir(parents=True, exist_ok=True)
report_path.write_text(json.dumps(data, indent=2, sort_keys=True) + "\n")
return report_path
def static_calibration_fields(config: TrainingConfig, model: nn.Module) -> dict[str, Any]:
forward_flops = estimate_forward_flops_per_item(model)
return {
"estimated_forward_flops_per_item": forward_flops,
"estimated_train_flops": estimate_training_compute(
steps=config.optim.steps,
batch_size=config.data.batch_size,
forward_flops_per_item=forward_flops,
),
}

View file

@ -6,18 +6,6 @@ from pathlib import Path
from typing import Any from typing import Any
_MODEL_TYPES = {
"mlp",
"film_fourier_mlp",
"film_fourier_inr",
"nerf_cfd_multires",
"deeponet_branch_trunk",
"point_context_perceiver",
"meshgraphnet_or_point_transformer_local",
"raster_fno_unet",
"siren_conditioned_inr",
}
@dataclass(frozen=True) @dataclass(frozen=True)
class RunConfig: class RunConfig:
name: str name: str
@ -33,19 +21,6 @@ class DataConfig:
test_cases: int test_cases: int
points_per_case: int points_per_case: int
batch_size: int batch_size: int
source: str
hf_repo_id: str | None
hf_repo_type: str
hf_path_prefix: str
cache_dir: Path | None
public_source_url: str | None = None
streaming_scratch_dir: Path | None = None
streaming_cache_max_bytes: int = 32 * 1024 * 1024 * 1024
streaming_cache_high_water_bytes: int = 28 * 1024 * 1024 * 1024
streaming_cache_low_water_bytes: int = 20 * 1024 * 1024 * 1024
streaming_queue_max_cases: int = 2
streaming_upload_processed: bool = False
streaming_upload_batch_size: int = 8
@dataclass(frozen=True) @dataclass(frozen=True)
@ -59,14 +34,6 @@ class ModelConfig:
condition_width: int condition_width: int
condition_depth: int condition_depth: int
condition_dim: int condition_dim: int
encoding_levels: int
features_per_level: int
context_points: int
latent_width: int
attention_depth: int
neighbors: int
grid_resolution: int
siren_omega0: float
@dataclass(frozen=True) @dataclass(frozen=True)
@ -109,20 +76,10 @@ class ObservabilityConfig:
backend: str backend: str
project: str project: str
entity: str | None entity: str | None
group: str | None
mode: str mode: str
tags: tuple[str, ...] tags: tuple[str, ...]
@dataclass(frozen=True)
class HuggingFaceConfig:
enabled: bool
repo_id: str
repo_type: str
path_prefix: str
private: bool
@dataclass(frozen=True) @dataclass(frozen=True)
class TrainingConfig: class TrainingConfig:
path: Path path: Path
@ -137,7 +94,6 @@ class TrainingConfig:
stability: StabilityConfig stability: StabilityConfig
precision: PrecisionConfig precision: PrecisionConfig
observability: ObservabilityConfig observability: ObservabilityConfig
huggingface: HuggingFaceConfig
_REQUIRED_SECTIONS = ("run", "data", "model", "optim", "device", "loss") _REQUIRED_SECTIONS = ("run", "data", "model", "optim", "device", "loss")
@ -185,12 +141,6 @@ def load_training_config(path: str | Path) -> TrainingConfig:
observability_raw = {} observability_raw = {}
if not isinstance(observability_raw, dict): if not isinstance(observability_raw, dict):
raise ValueError("Training config [observability] section must be a table") raise ValueError("Training config [observability] section must be a table")
huggingface_raw = raw.get("huggingface", {})
if huggingface_raw is None:
huggingface_raw = {}
if not isinstance(huggingface_raw, dict):
raise ValueError("Training config [huggingface] section must be a table")
run = RunConfig( run = RunConfig(
@ -198,24 +148,6 @@ def load_training_config(path: str | Path) -> TrainingConfig:
seed=_integer(run_raw, "seed", minimum=0), seed=_integer(run_raw, "seed", minimum=0),
artifact_dir=_path(run_raw, "artifact_dir"), artifact_dir=_path(run_raw, "artifact_dir"),
) )
streaming_cache_max_bytes = _integer(data_raw, "streaming_cache_max_bytes", minimum=1, default=32 * 1024 * 1024 * 1024)
streaming_high_water_bytes = _integer(
data_raw,
"streaming_cache_high_water_bytes",
minimum=1,
default=max(1, streaming_cache_max_bytes * 9 // 10),
)
streaming_low_water_bytes = _integer(
data_raw,
"streaming_cache_low_water_bytes",
minimum=1,
default=max(1, streaming_cache_max_bytes * 7 // 10),
)
if streaming_high_water_bytes > streaming_cache_max_bytes:
raise ValueError("data.streaming_cache_high_water_bytes must be <= data.streaming_cache_max_bytes")
if streaming_low_water_bytes >= streaming_high_water_bytes:
raise ValueError("data.streaming_cache_low_water_bytes must be < data.streaming_cache_high_water_bytes")
data = DataConfig( data = DataConfig(
root=_path(data_raw, "root"), root=_path(data_raw, "root"),
train_cases=_integer(data_raw, "train_cases", minimum=1), train_cases=_integer(data_raw, "train_cases", minimum=1),
@ -223,22 +155,9 @@ def load_training_config(path: str | Path) -> TrainingConfig:
test_cases=_integer(data_raw, "test_cases", minimum=0), test_cases=_integer(data_raw, "test_cases", minimum=0),
points_per_case=_integer(data_raw, "points_per_case", minimum=1), points_per_case=_integer(data_raw, "points_per_case", minimum=1),
batch_size=_integer(data_raw, "batch_size", minimum=1), batch_size=_integer(data_raw, "batch_size", minimum=1),
source=_choice(_string(data_raw, "source", default="local").lower(), {"local", "huggingface", "public_zip_streaming"}, "data.source"),
hf_repo_id=_optional_string(data_raw, "hf_repo_id"),
hf_repo_type=_choice(_string(data_raw, "hf_repo_type", default="dataset"), {"dataset"}, "data.hf_repo_type"),
hf_path_prefix=_string(data_raw, "hf_path_prefix", default=""),
cache_dir=_path(data_raw, "cache_dir") if "cache_dir" in data_raw else None,
public_source_url=_optional_string(data_raw, "public_source_url"),
streaming_scratch_dir=_path(data_raw, "streaming_scratch_dir") if "streaming_scratch_dir" in data_raw else None,
streaming_cache_max_bytes=streaming_cache_max_bytes,
streaming_cache_high_water_bytes=streaming_high_water_bytes,
streaming_cache_low_water_bytes=streaming_low_water_bytes,
streaming_queue_max_cases=_integer(data_raw, "streaming_queue_max_cases", minimum=1, default=2),
streaming_upload_processed=_boolean(data_raw, "streaming_upload_processed") if "streaming_upload_processed" in data_raw else False,
streaming_upload_batch_size=_integer(data_raw, "streaming_upload_batch_size", minimum=1, default=8),
) )
model = ModelConfig( model = ModelConfig(
type=_choice(_string(model_raw, "type"), _MODEL_TYPES, "model.type"), type=_choice(_string(model_raw, "type"), {"mlp", "film_fourier_mlp"}, "model.type"),
hidden_width=_integer(model_raw, "hidden_width", minimum=1), hidden_width=_integer(model_raw, "hidden_width", minimum=1),
depth=_integer(model_raw, "depth", minimum=1), depth=_integer(model_raw, "depth", minimum=1),
activation=_choice(_string(model_raw, "activation").lower(), {"gelu", "relu", "silu", "tanh"}, "model.activation"), activation=_choice(_string(model_raw, "activation").lower(), {"gelu", "relu", "silu", "tanh"}, "model.activation"),
@ -247,14 +166,6 @@ def load_training_config(path: str | Path) -> TrainingConfig:
condition_width=_integer(model_raw, "condition_width", minimum=1, default=_integer(model_raw, "hidden_width", minimum=1)), condition_width=_integer(model_raw, "condition_width", minimum=1, default=_integer(model_raw, "hidden_width", minimum=1)),
condition_depth=_integer(model_raw, "condition_depth", minimum=1, default=2), condition_depth=_integer(model_raw, "condition_depth", minimum=1, default=2),
condition_dim=_integer(model_raw, "condition_dim", minimum=1, default=_integer(model_raw, "hidden_width", minimum=1)), condition_dim=_integer(model_raw, "condition_dim", minimum=1, default=_integer(model_raw, "hidden_width", minimum=1)),
encoding_levels=_integer(model_raw, "encoding_levels", minimum=0, default=8),
features_per_level=_integer(model_raw, "features_per_level", minimum=1, default=2),
context_points=_integer(model_raw, "context_points", minimum=1, default=512),
latent_width=_integer(model_raw, "latent_width", minimum=1, default=_integer(model_raw, "hidden_width", minimum=1)),
attention_depth=_integer(model_raw, "attention_depth", minimum=1, default=2),
neighbors=_integer(model_raw, "neighbors", minimum=0, default=8),
grid_resolution=_integer(model_raw, "grid_resolution", minimum=2, default=32),
siren_omega0=_number(model_raw, "siren_omega0", minimum=0.0, exclusive_minimum=True, default=30.0),
) )
optim = OptimConfig( optim = OptimConfig(
lr=_number(optim_raw, "lr", minimum=0.0, exclusive_minimum=True), lr=_number(optim_raw, "lr", minimum=0.0, exclusive_minimum=True),
@ -279,18 +190,10 @@ def load_training_config(path: str | Path) -> TrainingConfig:
) )
observability = ObservabilityConfig( observability = ObservabilityConfig(
backend=_choice(_string(observability_raw, "backend", default="none").lower(), {"none", "wandb"}, "observability.backend"), backend=_choice(_string(observability_raw, "backend", default="none").lower(), {"none", "wandb"}, "observability.backend"),
project=_string(observability_raw, "project", default="airfRANS-model-sweep"), project=_string(observability_raw, "project", default="airfrans"),
entity=_optional_string(observability_raw, "entity"), entity=_optional_string(observability_raw, "entity"),
group=_optional_string(observability_raw, "group"),
mode=_choice(_string(observability_raw, "mode", default="online").lower(), {"online", "offline", "disabled"}, "observability.mode"), mode=_choice(_string(observability_raw, "mode", default="online").lower(), {"online", "offline", "disabled"}, "observability.mode"),
tags=_string_tuple(observability_raw, "tags", default=("airfrans",)), tags=_string_tuple(observability_raw, "tags", default=()),
)
huggingface = HuggingFaceConfig(
enabled=_boolean(huggingface_raw, "enabled") if "enabled" in huggingface_raw else False,
repo_id=_string(huggingface_raw, "repo_id", default="zacheryasc/airfrans-frontier-checkpoints"),
repo_type=_choice(_string(huggingface_raw, "repo_type", default="model"), {"model", "dataset", "space"}, "huggingface.repo_type"),
path_prefix=_string(huggingface_raw, "path_prefix", default="training_runs"),
private=_boolean(huggingface_raw, "private") if "private" in huggingface_raw else False,
) )
@ -311,7 +214,6 @@ def load_training_config(path: str | Path) -> TrainingConfig:
stability=stability, stability=stability,
precision=precision, precision=precision,
observability=observability, observability=observability,
huggingface=huggingface,
) )
@ -365,15 +267,8 @@ def _number(
*, *,
minimum: float | None = None, minimum: float | None = None,
exclusive_minimum: bool = False, exclusive_minimum: bool = False,
default: float | None = None,
) -> float: ) -> float:
if key not in section:
if default is None:
value = _required(section, key) value = _required(section, key)
else:
value = default
else:
value = section[key]
if isinstance(value, bool) or not isinstance(value, (int, float)): if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"Expected number for {key}") raise ValueError(f"Expected number for {key}")
result = float(value) result = float(value)

View file

@ -1,150 +0,0 @@
from __future__ import annotations
import hashlib
import json
import os
import time
from pathlib import Path
from typing import Any
from airfrans_frontier.training.config import DataConfig
def resolve_training_data_root(data: DataConfig) -> Path:
if data.source == "local":
return data.root
if data.source != "huggingface":
raise ValueError(f"Unsupported data source: {data.source}")
if not data.hf_repo_id:
raise ValueError("data.hf_repo_id is required when data.source = 'huggingface'")
try:
from huggingface_hub import snapshot_download
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required when data.source = 'huggingface'") from exc
token = _resolve_optional_token("HF_TOKEN")
prefix = data.hf_path_prefix.strip("/")
allow_patterns = [f"{prefix}/**"] if prefix else ["*.npz", "*.json", "*.jsonl", "*.txt"]
local_dir = data.cache_dir or data.root
local_dir.mkdir(parents=True, exist_ok=True)
downloaded = Path(
snapshot_download(
repo_id=data.hf_repo_id,
repo_type=data.hf_repo_type,
allow_patterns=allow_patterns,
local_dir=str(local_dir),
token=token,
)
)
resolved = downloaded / prefix if prefix else downloaded
if not resolved.is_dir():
raise FileNotFoundError(f"Downloaded Hugging Face data path is missing: {resolved}")
return resolved
def publish_processed_dataset(
*,
data_root: str | Path,
repo_id: str,
path_in_repo: str,
private: bool = False,
manifest_out: str | Path | None = None,
) -> dict[str, Any]:
root = Path(data_root).expanduser()
if not root.is_dir():
raise NotADirectoryError(f"Processed data path is not a directory: {root}")
files = sorted(path for path in root.rglob("*") if path.is_file() and not path.is_symlink())
npz_files = [path for path in files if path.suffix == ".npz"]
if not npz_files:
raise ValueError(f"No .npz simulation files found under: {root}")
path_prefix = path_in_repo.strip("/")
manifest = {
"repo_id": repo_id,
"repo_type": "dataset",
"repo_url": f"https://huggingface.co/datasets/{repo_id}",
"path_in_repo": path_prefix,
"created_at": time.time(),
"source_root": str(root),
"file_count": len(files),
"npz_file_count": len(npz_files),
"total_bytes": sum(path.stat().st_size for path in files),
"files": [
{
"path": str(path.relative_to(root)),
"bytes": path.stat().st_size,
"sha256": _sha256_file(path),
}
for path in files
],
}
manifest_path = Path(manifest_out).expanduser() if manifest_out is not None else root / "hf_dataset_manifest.json"
manifest_path.parent.mkdir(parents=True, exist_ok=True)
manifest_path.write_text(json.dumps(manifest, indent=2, sort_keys=True) + "\n")
try:
from huggingface_hub import HfApi
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required to publish processed data") from exc
token = _resolve_required_token("HF_TOKEN", "Hugging Face dataset publishing")
api = HfApi(token=token)
api.create_repo(repo_id=repo_id, repo_type="dataset", private=private, exist_ok=True)
commit = api.upload_folder(
repo_id=repo_id,
repo_type="dataset",
folder_path=str(root),
path_in_repo=path_prefix,
commit_message=f"Publish AirfRANS processed dataset {path_prefix or 'root'}",
)
manifest_repo_path = f"{path_prefix}/hf_dataset_manifest.json" if path_prefix else "hf_dataset_manifest.json"
api.upload_file(
repo_id=repo_id,
repo_type="dataset",
path_or_fileobj=str(manifest_path),
path_in_repo=manifest_repo_path,
commit_message=f"Add AirfRANS dataset manifest {path_prefix or 'root'}",
)
manifest["uploaded_manifest_path"] = manifest_repo_path
manifest["commit"] = {
"commit_url": getattr(commit, "commit_url", None),
"commit_hash": getattr(commit, "oid", None) or getattr(commit, "commit_hash", None),
"pr_url": getattr(commit, "pr_url", None),
}
manifest_path.write_text(json.dumps(manifest, indent=2, sort_keys=True) + "\n")
return manifest
def _resolve_optional_token(name: str) -> str | None:
try:
return _resolve_required_token(name, "optional Hugging Face access")
except RuntimeError:
return None
def _resolve_required_token(name: str, purpose: str) -> str:
value = os.environ.get(name)
if value and value.strip():
return value.strip()
for path in (Path(name), Path(".env") / name):
if path.is_file():
value = path.read_text().strip()
if value:
return value
env_file = Path(".env")
if env_file.is_file():
for raw_line in env_file.read_text().splitlines():
line = raw_line.strip()
if not line or line.startswith("#") or "=" not in line:
continue
key, value = line.split("=", 1)
if key.strip() == name:
value = value.strip().strip("\"'")
if value:
return value
raise RuntimeError(f"{name} env var or local secret file is required for {purpose}")
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as file:
for chunk in iter(lambda: file.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()

View file

@ -1,44 +0,0 @@
from __future__ import annotations
import platform
import subprocess
import sys
from typing import Any
def environment_manifest() -> dict[str, Any]:
manifest: dict[str, Any] = {
"python": sys.version,
"platform": platform.platform(),
"executable": sys.executable,
}
try:
import torch
manifest.update(
{
"torch_version": torch.__version__,
"cuda_available": torch.cuda.is_available(),
"cuda_version": torch.version.cuda,
"gpu_name": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
}
)
except Exception as exc:
manifest["torch_error"] = repr(exc)
try:
result = subprocess.run(
["nvidia-smi", "--query-gpu=name,memory.total,driver_version", "--format=csv,noheader"],
text=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
timeout=15,
check=False,
)
manifest["nvidia_smi"] = {
"returncode": result.returncode,
"stdout": result.stdout.strip(),
"stderr": result.stderr.strip(),
}
except OSError as exc:
manifest["nvidia_smi"] = {"error": str(exc)}
return manifest

View file

@ -1,357 +0,0 @@
from __future__ import annotations
import hashlib
import re
import json
import os
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Iterable
from airfrans_frontier.training.config import TrainingConfig
@dataclass
class UploadRecord:
local_path: str
repo_path: str
bytes: int
sha256: str
uploaded_at: float
@dataclass
class UploadManifest:
enabled: bool
repo_id: str | None
repo_type: str | None
repo_url: str | None
path_in_repo: str | None
uploaded_paths: list[str] = field(default_factory=list)
uploaded_files: list[dict[str, Any]] = field(default_factory=list)
commits: list[dict[str, Any]] = field(default_factory=list)
suppressed_uploads: list[dict[str, Any]] = field(default_factory=list)
last_error: str | None = None
rate_limit_until: float | None = None
rate_limit_retry_after_seconds: float | None = None
training_success: bool | None = None
publication_complete: bool = False
publication_status: str = "disabled"
finalized_at: float | None = None
class HfArtifactUploader:
def __init__(
self,
*,
enabled: bool,
run_dir: Path,
repo_id: str | None = None,
repo_type: str | None = None,
path_in_repo: str | None = None,
private: bool = False,
max_rate_limit_sleep_seconds: float = 300.0,
) -> None:
self.enabled = enabled
self.run_dir = run_dir
self.repo_id = repo_id
self.repo_type = repo_type
self.path_in_repo = path_in_repo.strip("/") if path_in_repo else None
self.private = private
self.max_rate_limit_sleep_seconds = max_rate_limit_sleep_seconds
self._api: Any | None = None
self._manifest = UploadManifest(
enabled=enabled,
repo_id=repo_id,
repo_type=repo_type,
repo_url=f"https://huggingface.co/{repo_id}" if repo_id else None,
path_in_repo=self.path_in_repo,
)
self.write_manifest()
@classmethod
def from_config(cls, config: TrainingConfig, *, run_dir: Path, run_id: str) -> HfArtifactUploader:
path_prefix = config.huggingface.path_prefix.strip("/")
path_parts = [part for part in (path_prefix, config.model.type, run_id) if part]
return cls(
enabled=config.huggingface.enabled,
run_dir=run_dir,
repo_id=config.huggingface.repo_id,
repo_type=config.huggingface.repo_type,
path_in_repo="/".join(path_parts),
private=config.huggingface.private,
)
@property
def repo_url(self) -> str | None:
return self._manifest.repo_url
@property
def publication_status(self) -> str:
self._refresh_publication_status()
return self._manifest.publication_status
@property
def publication_complete(self) -> bool:
self._refresh_publication_status()
return self._manifest.publication_complete
def finalize(self, *, training_success: bool) -> dict[str, Any]:
self._manifest.training_success = bool(training_success)
self._manifest.finalized_at = time.time()
self._refresh_publication_status()
self.write_manifest()
return self.final_report()
def final_report(self) -> dict[str, Any]:
self._refresh_publication_status()
return {
"hf_publication_status": self._manifest.publication_status,
"hf_publication_complete": self._manifest.publication_complete,
"hf_training_success": self._manifest.training_success,
"hf_rate_limit_until": self._manifest.rate_limit_until,
"hf_last_error": self._manifest.last_error,
}
def upload_files(self, names: Iterable[str], *, commit_message: str) -> dict[str, Any]:
names = tuple(dict.fromkeys(names))
if not self.enabled:
return {"enabled": False, "uploaded": [], "missing": [], "rate_limited": False}
missing = [name for name in names if not (self.run_dir / name).is_file()]
if missing:
raise FileNotFoundError(f"Cannot upload missing Hugging Face artifacts: {', '.join(missing)}")
suppressed = self._suppress_if_rate_limited(names, commit_message=commit_message)
if suppressed is not None:
return suppressed
api = self._ensure_api()
paths: list[tuple[Path, str]] = []
for name in names:
local_path = self.run_dir / name
repo_path = f"{self.path_in_repo}/{name}" if self.path_in_repo else name
paths.append((local_path, repo_path))
uploaded = [repo_path for _, repo_path in paths]
attempts = 0
while True:
try:
from huggingface_hub import CommitOperationAdd
operations = [
CommitOperationAdd(path_in_repo=repo_path, path_or_fileobj=str(local_path))
for local_path, repo_path in paths
]
commit = api.create_commit(
repo_id=self.repo_id,
repo_type=self.repo_type,
operations=operations,
commit_message=commit_message,
)
uploaded_at = time.time()
for local_path, repo_path in paths:
record = UploadRecord(
local_path=str(local_path),
repo_path=repo_path,
bytes=local_path.stat().st_size,
sha256=_sha256_file(local_path),
uploaded_at=uploaded_at,
)
self._manifest.uploaded_paths.append(repo_path)
self._manifest.uploaded_files.append(record.__dict__)
self._manifest.commits.append(_commit_payload(commit))
self._manifest.uploaded_paths = sorted(set(self._manifest.uploaded_paths))
self._manifest.last_error = None
self._manifest.rate_limit_until = None
self._manifest.rate_limit_retry_after_seconds = None
self.write_manifest()
return {"enabled": True, "uploaded": uploaded, "missing": [], "rate_limited": False}
except Exception as exc:
retry_after = _retry_after_seconds(exc)
if retry_after is not None:
self._record_rate_limit(exc, retry_after, names=names, commit_message=commit_message)
if attempts == 0 and retry_after <= self.max_rate_limit_sleep_seconds:
attempts += 1
time.sleep(max(0.0, retry_after))
continue
self._manifest.last_error = str(exc)
self.write_manifest()
raise
def _suppress_if_rate_limited(self, names: tuple[str, ...], *, commit_message: str) -> dict[str, Any] | None:
until = self._manifest.rate_limit_until
now = time.time()
if until is None or now >= until:
return None
record = {
"names": list(names),
"commit_message": commit_message,
"suppressed_at": now,
"rate_limit_until": until,
}
self._manifest.suppressed_uploads.append(record)
self._manifest.last_error = f"HF upload suppressed until {until:.3f} after rate limiting"
self.write_manifest()
return {"enabled": True, "uploaded": [], "missing": [], "rate_limited": True, "suppressed_until": until}
def _record_rate_limit(self, exc: Exception, retry_after: float, *, names: tuple[str, ...], commit_message: str) -> None:
now = time.time()
until = now + retry_after
self._manifest.rate_limit_until = max(self._manifest.rate_limit_until or 0.0, until)
self._manifest.rate_limit_retry_after_seconds = retry_after
self._manifest.last_error = str(exc)
self._manifest.suppressed_uploads.append(
{
"names": list(names),
"commit_message": commit_message,
"rate_limited_at": now,
"rate_limit_until": self._manifest.rate_limit_until,
"retry_after_seconds": retry_after,
}
)
self.write_manifest()
def write_manifest(self) -> Path:
self._refresh_publication_status()
path = self.run_dir / "hf_upload_manifest.json"
path.write_text(json.dumps(self._manifest.__dict__, indent=2, sort_keys=True) + "\n")
return path
def _refresh_publication_status(self) -> None:
if not self.enabled:
self._manifest.publication_status = "disabled"
self._manifest.publication_complete = False
return
if self._manifest.training_success is False:
self._manifest.publication_status = "training_failed"
self._manifest.publication_complete = False
return
incomplete = self._manifest.last_error is not None or self._manifest.rate_limit_until is not None
if incomplete:
self._manifest.publication_status = (
"training_succeeded_hf_incomplete" if self._manifest.training_success is True else "hf_publication_incomplete"
)
self._manifest.publication_complete = False
return
if self._manifest.training_success is True:
self._manifest.publication_status = "hf_publication_succeeded"
self._manifest.publication_complete = True
return
self._manifest.publication_status = "in_progress"
self._manifest.publication_complete = False
def _ensure_api(self) -> Any:
if self._api is not None:
return self._api
try:
from huggingface_hub import HfApi
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required when [huggingface].enabled = true") from exc
token = _resolve_secret("HF_TOKEN", "Hugging Face checkpoint uploads")
api = HfApi(token=token)
assert self.repo_id is not None
assert self.repo_type is not None
api.create_repo(repo_id=self.repo_id, repo_type=self.repo_type, private=self.private, exist_ok=True)
self._api = api
return api
def resolve_resume_checkpoint(resume_path: str | Path | None) -> tuple[Path | None, dict[str, Any]]:
if resume_path is None:
return None, {"resume_source": None, "resume_downloaded": False, "resume_downloaded_path": None}
raw = str(resume_path)
if not raw.startswith("hf://"):
return Path(resume_path).expanduser(), {"resume_source": raw, "resume_downloaded": False, "resume_downloaded_path": None}
repo_id, filename = _parse_hf_checkpoint_uri(raw)
try:
from huggingface_hub import hf_hub_download
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required to resume from hf:// checkpoints") from exc
token = _resolve_secret("HF_TOKEN", "Hugging Face checkpoint download")
cache_dir = Path(".airfrans_hf_resume") / hashlib.sha256(raw.encode()).hexdigest()[:16]
cache_dir.mkdir(parents=True, exist_ok=True)
downloaded = Path(
hf_hub_download(
repo_id=repo_id,
repo_type="model",
filename=filename,
token=token,
local_dir=str(cache_dir),
)
)
return downloaded, {"resume_source": raw, "resume_downloaded": True, "resume_downloaded_path": str(downloaded)}
def _parse_hf_checkpoint_uri(uri: str) -> tuple[str, str]:
rest = uri.removeprefix("hf://")
parts = rest.split("/")
if len(parts) < 3:
raise ValueError("HF checkpoint URI must be hf://namespace/repo/path/to/checkpoint.pt")
repo_id = "/".join(parts[:2])
filename = "/".join(parts[2:])
if not filename:
raise ValueError("HF checkpoint URI is missing checkpoint path")
return repo_id, filename
def _resolve_secret(name: str, purpose: str) -> str:
value = os.environ.get(name)
if value and value.strip():
return value.strip()
for path in (Path(name), Path(".env") / name):
if path.is_file():
value = path.read_text().strip()
if value:
return value
env_file = Path(".env")
if env_file.is_file():
for raw_line in env_file.read_text().splitlines():
line = raw_line.strip()
if not line or line.startswith("#") or "=" not in line:
continue
key, value = line.split("=", 1)
if key.strip() == name:
value = value.strip().strip("\"'")
if value:
return value
raise RuntimeError(f"{name} env var or local secret file is required for {purpose}")
def _retry_after_seconds(exc: Exception) -> float | None:
response = getattr(exc, "response", None)
headers = getattr(response, "headers", None)
if headers is not None:
raw = headers.get("Retry-After") or headers.get("retry-after")
if raw is not None:
parsed = _parse_retry_after(raw)
if parsed is not None:
return parsed
match = re.search(r"Retry after\s+(\d+(?:\.\d+)?)\s+seconds", str(exc), flags=re.IGNORECASE)
if match:
return float(match.group(1))
if "rate limit" not in str(exc).lower() and "too many requests" not in str(exc).lower():
return None
return 300.0
def _parse_retry_after(value: object) -> float | None:
try:
seconds = float(str(value).strip())
except ValueError:
return None
if seconds < 0:
return None
return seconds
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as file:
for chunk in iter(lambda: file.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _commit_payload(commit: Any) -> dict[str, Any]:
return {
"commit_url": getattr(commit, "commit_url", None),
"commit_hash": getattr(commit, "oid", None) or getattr(commit, "commit_hash", None),
"pr_url": getattr(commit, "pr_url", None),
}

File diff suppressed because it is too large Load diff

View file

@ -18,9 +18,6 @@ class TrainingObserver:
def update_summary(self, metrics: Mapping[str, Any]) -> None: def update_summary(self, metrics: Mapping[str, Any]) -> None:
return None return None
def update_config(self, values: Mapping[str, Any]) -> None:
return None
def finish(self, *, exit_code: int = 0) -> None: def finish(self, *, exit_code: int = 0) -> None:
return None return None
@ -50,9 +47,6 @@ class WandbObserver(TrainingObserver):
for key, value in _json_safe(dict(metrics)).items(): for key, value in _json_safe(dict(metrics)).items():
self._run.summary[key] = value self._run.summary[key] = value
def update_config(self, values: Mapping[str, Any]) -> None:
self._run.config.update(_json_safe(dict(values)), allow_val_change=True)
def finish(self, *, exit_code: int = 0) -> None: def finish(self, *, exit_code: int = 0) -> None:
self._wandb.finish(exit_code=exit_code) self._wandb.finish(exit_code=exit_code)
@ -72,7 +66,6 @@ def start_observer(config: TrainingConfig, *, run_dir: Path) -> TrainingObserver
run = wandb.init( run = wandb.init(
entity=observability.entity, entity=observability.entity,
project=observability.project, project=observability.project,
group=observability.group,
name=config.run.name, name=config.run.name,
tags=list(observability.tags), tags=list(observability.tags),
mode=observability.mode, mode=observability.mode,

View file

@ -1,167 +0,0 @@
from __future__ import annotations
import json
import tempfile
from pathlib import Path
from typing import Any, Iterable
import numpy as np
import torch
from airfrans_frontier.training.config import load_training_config
from airfrans_frontier.training.loop import train
MODEL_FAMILIES = (
"film_fourier_inr",
"nerf_cfd_multires",
"deeponet_branch_trunk",
"point_context_perceiver",
"meshgraphnet_or_point_transformer_local",
"raster_fno_unet",
"siren_conditioned_inr",
)
FEATURE_NAMES = np.array(["x", "y", "sdf", "u_inf", "log_re", "aoa_deg", "aoa_sin", "aoa_cos"], dtype="U16")
TARGET_NAMES = np.array(["velocity_x", "velocity_y", "pressure", "turbulent_viscosity"], dtype="U32")
def run_model_sanity(
*,
artifact_dir: str | Path,
device_type: str = "auto",
families: Iterable[str] = MODEL_FAMILIES,
steps: int = 80,
) -> dict[str, Any]:
output_dir = Path(artifact_dir)
output_dir.mkdir(parents=True, exist_ok=True)
selected = tuple(families)
results: dict[str, Any] = {
"device_requested": device_type,
"cuda_available": torch.cuda.is_available(),
"families": {},
}
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
data_root = tmp_path / "toy_data"
_write_toy_dataset(data_root)
for family in selected:
config_path = tmp_path / f"{family}.toml"
family_artifacts = output_dir / family
config_path.write_text(_config_text(family, data_root=data_root, artifact_dir=family_artifacts, device_type=device_type, steps=steps))
result = train(load_training_config(config_path))
final_metrics = result.final_metrics
initial = float(final_metrics["initial_train_loss"])
final = float(final_metrics["train_loss"])
decreased = final < initial
results["families"][family] = {
"run_dir": str(result.run_dir),
"initial_train_loss": initial,
"final_train_loss": final,
"loss_decreased": decreased,
"device": final_metrics.get("device"),
"parameter_count": final_metrics.get("parameter_count"),
"points_per_sec": final_metrics.get("points_per_sec"),
"step_time_seconds": final_metrics.get("step_time_seconds"),
"validation_runtime_seconds": final_metrics.get("validation_runtime_seconds"),
"checkpoint_latest_bytes": final_metrics.get("checkpoint_latest_bytes"),
"checkpoint_best_bytes": final_metrics.get("checkpoint_best_bytes"),
"checkpoint_final_bytes": final_metrics.get("checkpoint_final_bytes"),
"gpu_memory_peak_allocated_mb": final_metrics.get("gpu_memory_peak_allocated_mb"),
"estimated_forward_flops_per_item": final_metrics.get("estimated_forward_flops_per_item"),
"estimated_train_flops": final_metrics.get("estimated_train_flops"),
"target_context_policy": final_metrics.get("target_context_policy"),
"locality_protocol": final_metrics.get("locality_protocol"),
"raster_protocol": final_metrics.get("raster_protocol"),
}
if not decreased:
raise RuntimeError(f"Toy sanity loss did not decrease for {family}: initial={initial}, final={final}")
results["ok"] = True
report_path = output_dir / "model_sanity_results.json"
report_path.write_text(json.dumps(results, indent=2, sort_keys=True) + "\n")
return results
def _write_toy_dataset(root: Path, *, cases: int = 4, points: int = 64) -> None:
root.mkdir(parents=True, exist_ok=True)
rng = np.random.default_rng(8675309)
for case_index in range(cases):
x = rng.uniform(-1.0, 1.0, size=points).astype(np.float32)
y = rng.uniform(-1.0, 1.0, size=points).astype(np.float32)
sdf = (np.sqrt(x * x + y * y) - 0.35).astype(np.float32)
aoa = np.float32(-6.0 + 4.0 * case_index)
u_inf = np.float32(20.0 + 2.0 * case_index)
log_re = np.log(u_inf / np.float32(1.5e-5)).astype(np.float32)
condition = np.tile(
np.array([u_inf / 40.0, log_re / 20.0, aoa / 10.0, np.sin(np.deg2rad(aoa)), np.cos(np.deg2rad(aoa))], dtype=np.float32),
(points, 1),
)
features = np.concatenate((np.stack((x, y, sdf), axis=1), condition), axis=1).astype(np.float32)
targets = np.stack(
(
0.35 * x + 0.10 * y + 0.04 * aoa,
-0.25 * y + 0.02 * u_inf / 40.0,
x * y + 0.05 * sdf,
sdf * sdf + 0.03 * np.sin(np.deg2rad(aoa)) + 0.01 * x,
),
axis=1,
).astype(np.float32)
np.savez(root / f"case_{case_index:02d}.npz", features=features, targets=targets, feature_names=FEATURE_NAMES, target_names=TARGET_NAMES)
def _config_text(family: str, *, data_root: Path, artifact_dir: Path, device_type: str, steps: int) -> str:
return f"""
[run]
name = "sanity_{family}"
seed = 7
artifact_dir = "{artifact_dir}"
[data]
root = "{data_root}"
train_cases = 2
val_cases = 1
test_cases = 1
points_per_case = 64
batch_size = 32
[model]
type = "{family}"
hidden_width = 24
depth = 2
activation = "gelu"
coordinate_features = ["x", "y", "sdf"]
fourier_scales = [1.0, 2.0]
condition_width = 24
condition_depth = 2
condition_dim = 24
encoding_levels = 3
features_per_level = 2
context_points = 16
latent_width = 24
attention_depth = 2
neighbors = 4
grid_resolution = 8
siren_omega0 = 10.0
[optim]
lr = 0.01
weight_decay = 0.0
steps = {steps}
log_interval = {max(1, steps // 4)}
[device]
type = "{device_type}"
allow_cpu_fallback = true
benchmark_kernels = false
[loss]
type = "normalized_mse"
[checkpoint]
interval_seconds = 0
[observability]
backend = "none"
[huggingface]
enabled = false
""".strip() + "\n"

File diff suppressed because it is too large Load diff

View file

@ -1,98 +0,0 @@
from __future__ import annotations
import json
import os
import sys
import tempfile
import types
import unittest
from pathlib import Path
from unittest.mock import patch
from airfrans_frontier.runtime import remove_pythonpath_entries
from airfrans_frontier.training.config import DataConfig
from airfrans_frontier.training.data_sources import publish_processed_dataset, resolve_training_data_root
remove_pythonpath_entries()
import numpy as np
class DataSourceTests(unittest.TestCase):
def test_huggingface_source_downloads_prefix_to_cache(self) -> None:
calls: list[dict[str, object]] = []
def fake_snapshot_download(**kwargs):
calls.append(kwargs)
local_dir = Path(str(kwargs["local_dir"]))
target = local_dir / "processed" / "full"
target.mkdir(parents=True)
return str(local_dir)
fake_module = types.SimpleNamespace(snapshot_download=fake_snapshot_download)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {}, clear=False):
tmp_path = Path(tmp)
config = DataConfig(
root=tmp_path / "configured-root",
train_cases=1,
val_cases=0,
test_cases=0,
points_per_case=1,
batch_size=1,
source="huggingface",
hf_repo_id="owner/airfrans-processed",
hf_repo_type="dataset",
hf_path_prefix="processed/full",
cache_dir=tmp_path / "cache",
)
resolved = resolve_training_data_root(config)
self.assertEqual(resolved, tmp_path / "cache" / "processed" / "full")
self.assertEqual(calls[0]["repo_id"], "owner/airfrans-processed")
self.assertEqual(calls[0]["repo_type"], "dataset")
self.assertEqual(calls[0]["allow_patterns"], ["processed/full/**"])
def test_publish_processed_dataset_uploads_folder_and_manifest(self) -> None:
created: list[tuple[str, str, bool]] = []
uploaded_folders: list[tuple[str, str]] = []
uploaded_files: list[str] = []
class FakeApi:
def __init__(self, token: str) -> None:
self.token = token
def create_repo(self, *, repo_id: str, repo_type: str, private: bool, exist_ok: bool) -> None:
created.append((repo_id, repo_type, private))
def upload_folder(self, *, repo_id: str, repo_type: str, folder_path: str, path_in_repo: str, commit_message: str):
uploaded_folders.append((folder_path, path_in_repo))
return types.SimpleNamespace(commit_url="https://huggingface.co/datasets/owner/repo/commit/abc", oid="abc")
def upload_file(self, *, repo_id: str, repo_type: str, path_or_fileobj: str, path_in_repo: str, commit_message: str):
uploaded_files.append(path_in_repo)
return types.SimpleNamespace(commit_url="https://huggingface.co/datasets/owner/repo/commit/def", oid="def")
fake_module = types.SimpleNamespace(HfApi=FakeApi)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}):
root = Path(tmp) / "processed"
root.mkdir()
np.savez(root / "case_00.npz", features=np.zeros((2, 2), dtype=np.float32), targets=np.zeros((2, 1), dtype=np.float32))
manifest_path = Path(tmp) / "manifest.json"
manifest = publish_processed_dataset(
data_root=root,
repo_id="owner/repo",
path_in_repo="processed/full",
manifest_out=manifest_path,
)
self.assertEqual(created, [("owner/repo", "dataset", False)])
self.assertEqual(uploaded_folders, [(str(root), "processed/full")])
self.assertEqual(uploaded_files, ["processed/full/hf_dataset_manifest.json"])
self.assertEqual(manifest["npz_file_count"], 1)
self.assertTrue(manifest_path.is_file())
self.assertEqual(json.loads(manifest_path.read_text())["uploaded_manifest_path"], "processed/full/hf_dataset_manifest.json")
if __name__ == "__main__":
unittest.main()

View file

@ -1,120 +0,0 @@
from __future__ import annotations
import json
import os
import sys
import tempfile
import types
import unittest
from pathlib import Path
from unittest.mock import Mock, patch
from airfrans_frontier.training.hf_upload import HfArtifactUploader, resolve_resume_checkpoint
class HuggingFaceUploadTests(unittest.TestCase):
def test_uploader_creates_repo_uploads_file_and_writes_manifest(self) -> None:
created: list[tuple[str, str, bool]] = []
committed: list[tuple[str, str, tuple[str, ...], str]] = []
class FakeCommitOperationAdd:
def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None:
self.path_in_repo = path_in_repo
self.path_or_fileobj = path_or_fileobj
class FakeApi:
def __init__(self, token: str) -> None:
self.token = token
def create_repo(self, *, repo_id: str, repo_type: str, private: bool, exist_ok: bool) -> None:
created.append((repo_id, repo_type, private))
def create_commit(self, *, repo_id: str, repo_type: str, operations: list[FakeCommitOperationAdd], commit_message: str):
committed.append((repo_id, repo_type, tuple(operation.path_in_repo for operation in operations), commit_message))
return types.SimpleNamespace(commit_url="https://huggingface.co/repo/commit/abc", oid="abc")
fake_module = types.SimpleNamespace(HfApi=FakeApi, CommitOperationAdd=FakeCommitOperationAdd)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}):
root = Path(tmp)
(root / "checkpoint_latest.pt").write_bytes(b"checkpoint")
uploader = HfArtifactUploader(
enabled=True,
run_dir=root,
repo_id="owner/repo",
repo_type="model",
path_in_repo="runs/model/run-1",
private=False,
)
result = uploader.upload_files(("checkpoint_latest.pt",), commit_message="upload checkpoint")
self.assertEqual(result["uploaded"], ["runs/model/run-1/checkpoint_latest.pt"])
self.assertEqual(created, [("owner/repo", "model", False)])
self.assertEqual(committed, [("owner/repo", "model", ("runs/model/run-1/checkpoint_latest.pt",), "upload checkpoint")])
manifest = json.loads((root / "hf_upload_manifest.json").read_text())
self.assertTrue(manifest["enabled"])
self.assertEqual(manifest["repo_id"], "owner/repo")
self.assertIn("runs/model/run-1/checkpoint_latest.pt", manifest["uploaded_paths"])
def test_uploader_suppresses_uploads_after_hf_retry_after_limit(self) -> None:
class FakeRateLimitError(RuntimeError):
def __init__(self) -> None:
super().__init__("429 Too Many Requests: Retry after 600 seconds")
self.response = types.SimpleNamespace(headers={"Retry-After": "600"})
class FakeCommitOperationAdd:
def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None:
self.path_in_repo = path_in_repo
self.path_or_fileobj = path_or_fileobj
fake_module = types.SimpleNamespace(CommitOperationAdd=FakeCommitOperationAdd)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}):
root = Path(tmp)
(root / "metrics.jsonl").write_text("{}\n")
uploader = HfArtifactUploader(
enabled=True,
run_dir=root,
repo_id="owner/repo",
repo_type="model",
path_in_repo="runs/model/run-1",
private=False,
max_rate_limit_sleep_seconds=0,
)
fake_api = types.SimpleNamespace(create_commit=Mock(side_effect=FakeRateLimitError()))
uploader._api = fake_api
with self.assertRaises(FakeRateLimitError):
uploader.upload_files(("metrics.jsonl",), commit_message="first")
suppressed = uploader.upload_files(("metrics.jsonl",), commit_message="second")
self.assertTrue(suppressed["rate_limited"])
self.assertEqual(fake_api.create_commit.call_count, 1)
manifest = json.loads((root / "hf_upload_manifest.json").read_text())
self.assertGreater(manifest["rate_limit_until"], 0)
self.assertEqual(manifest["rate_limit_retry_after_seconds"], 600.0)
self.assertEqual(len(manifest["suppressed_uploads"]), 2)
def test_resolve_resume_checkpoint_downloads_hf_uri(self) -> None:
calls: list[tuple[str, str]] = []
def fake_download(*, repo_id: str, repo_type: str, filename: str, token: str, local_dir: str) -> str:
calls.append((repo_id, filename))
path = Path(local_dir) / filename
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(b"checkpoint")
return str(path)
fake_module = types.SimpleNamespace(hf_hub_download=fake_download)
with patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}):
path, info = resolve_resume_checkpoint("hf://owner/repo/runs/model/checkpoint_latest.pt")
self.assertIsNotNone(path)
assert path is not None
self.assertTrue(path.is_file())
self.assertEqual(calls, [("owner/repo", "runs/model/checkpoint_latest.pt")])
self.assertTrue(info["resume_downloaded"])
self.assertEqual(info["resume_source"], "hf://owner/repo/runs/model/checkpoint_latest.pt")
if __name__ == "__main__":
unittest.main()

View file

@ -8,15 +8,7 @@ remove_pythonpath_entries()
import torch import torch
from airfrans_frontier.models import ( from airfrans_frontier.models import PointwiseMLP
DeepONetBranchTrunk,
LocalPointTransformer,
NeRFCFDMultiRes,
PointContextPerceiver,
PointwiseMLP,
RasterFNOUNet,
SirenConditionedINR,
)
class PointwiseMLPTests(unittest.TestCase): class PointwiseMLPTests(unittest.TestCase):
@ -28,79 +20,6 @@ class PointwiseMLPTests(unittest.TestCase):
self.assertEqual(tuple(output.shape), (7, 4)) self.assertEqual(tuple(output.shape), (7, 4))
def test_frontier_models_return_batch_by_target_dim(self) -> None:
feature_names = ("x", "y", "sdf", "u_inf", "log_re", "aoa_deg")
batch = torch.randn(8, len(feature_names))
models = [
NeRFCFDMultiRes(
feature_names=feature_names,
output_dim=4,
coordinate_features=("x", "y", "sdf"),
encoding_levels=2,
hidden_width=16,
depth=2,
condition_width=12,
condition_depth=2,
activation="gelu",
),
DeepONetBranchTrunk(
feature_names=feature_names,
output_dim=4,
coordinate_features=("x", "y", "sdf"),
fourier_scales=(1.0, 2.0),
hidden_width=16,
depth=2,
condition_width=12,
condition_depth=2,
activation="gelu",
),
PointContextPerceiver(
input_dim=len(feature_names),
output_dim=4,
hidden_width=16,
latent_width=12,
context_points=6,
attention_depth=2,
activation="gelu",
),
LocalPointTransformer(
feature_names=feature_names,
output_dim=4,
coordinate_features=("x", "y", "sdf"),
hidden_width=16,
depth=2,
neighbors=3,
activation="gelu",
),
RasterFNOUNet(
feature_names=feature_names,
output_dim=4,
coordinate_features=("x", "y", "sdf"),
grid_resolution=4,
hidden_width=16,
depth=2,
condition_width=12,
condition_depth=2,
activation="gelu",
),
SirenConditionedINR(
feature_names=feature_names,
output_dim=4,
coordinate_features=("x", "y", "sdf"),
hidden_width=16,
depth=2,
condition_width=12,
condition_depth=2,
omega0=10.0,
activation="gelu",
),
]
for model in models:
with self.subTest(model=type(model).__name__):
output = model(batch)
self.assertEqual(tuple(output.shape), (8, 4))
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()

View file

@ -1,279 +0,0 @@
from __future__ import annotations
import json
import shutil
import sys
import tempfile
import types
import unittest
from pathlib import Path
from unittest.mock import Mock, patch
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
from airfrans_frontier.remote.cleanup import reconcile_cleanup
from airfrans_frontier.remote.collection import ARTIFACT_COLLECTION_REPORT, collect_artifact_paths, required_collection_failures
from airfrans_frontier.remote.config import load_remote_run_config
from airfrans_frontier.remote.launch_group import LaunchGroupScheduler, LaunchRunSpec
from airfrans_frontier.remote.selection import require_fresh_selection, selection_freshness_report
from airfrans_frontier.remote.skypilot import render_skypilot_yaml
from airfrans_frontier.remote.vast import VastOffer, choose_offer
from airfrans_frontier.training.hf_upload import HfArtifactUploader
from airfrans_frontier.training.streaming_data import StreamingEventRecorder
class LaunchGroupSchedulingTests(unittest.TestCase):
def test_healthy_runs_release_fragile_launch_capacity_without_serializing_training(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
state_path = Path(tmp) / "launch_state.json"
scheduler = LaunchGroupScheduler(
[
LaunchRunSpec("run-a", "configs/a.toml"),
LaunchRunSpec("run-b", "configs/b.toml"),
LaunchRunSpec("run-c", "configs/c.toml"),
],
max_active=3,
max_fragile=1,
state_path=state_path,
group_id="group-local",
)
self.assertTrue(scheduler.try_start("run-a", selected_offer_id=101, selected_host_id=11))
self.assertFalse(scheduler.try_start("run-b", selected_offer_id=102, selected_host_id=12))
self.assertEqual(scheduler.capacity_snapshot()["fragile"], 1)
scheduler.mark_training_healthy("run-a")
self.assertTrue(scheduler.try_start("run-b", selected_offer_id=102, selected_host_id=12))
self.assertEqual(scheduler.capacity_snapshot()["active"], 2)
self.assertEqual(scheduler.capacity_snapshot()["fragile"], 1)
payload = json.loads(state_path.read_text())
self.assertEqual(payload["launch_group_id"], "group-local")
self.assertEqual(payload["healthy_runs"], ["run-a"])
self.assertEqual(payload["running_runs"], ["run-b"])
self.assertEqual(payload["runs"]["run-b"]["selected_host_id"], 12)
self.assertIn("capacity_blocked", [event["event"] for event in payload["events"]])
class HostAntiCollisionTests(unittest.TestCase):
def test_active_launches_avoid_duplicate_hosts_unless_allowed(self) -> None:
scheduler = LaunchGroupScheduler(
[LaunchRunSpec("run-a", "a.toml"), LaunchRunSpec("run-b", "b.toml")],
max_active=2,
max_fragile=2,
)
self.assertTrue(scheduler.try_start("run-a", selected_offer_id=1, selected_host_id=9))
self.assertFalse(scheduler.try_start("run-b", selected_offer_id=2, selected_host_id=9))
self.assertEqual(scheduler.to_payload()["runs"]["run-b"]["blocked_reason"], "host_collision")
allowed = LaunchGroupScheduler(
[LaunchRunSpec("run-a", "a.toml"), LaunchRunSpec("run-b", "b.toml")],
max_active=2,
max_fragile=2,
allow_duplicate_hosts=True,
)
self.assertTrue(allowed.try_start("run-a", selected_offer_id=1, selected_host_id=9))
self.assertTrue(allowed.try_start("run-b", selected_offer_id=2, selected_host_id=9))
def test_offer_selection_skips_reserved_active_hosts(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml")
result = choose_offer(
[offer(10, price=0.20, host=1), offer(11, price=0.22, host=2)],
config,
query={"test": True},
reserved_host_ids=(1,),
)
self.assertEqual(result.selected_offer.host_id, 2)
self.assertEqual(result.policy["reserved_host_ids"], [1])
class SelectionFreshnessTests(unittest.TestCase):
def test_selection_artifacts_record_and_enforce_freshness(self) -> None:
fresh = {"selected_offer_id": 1, "created_at": 1000.0}
report = selection_freshness_report(fresh, max_age_seconds=60, now=1020.0)
self.assertTrue(report["is_fresh"])
self.assertEqual(report["age_seconds"], 20.0)
stale = {"selected_offer_id": 1, "created_at": 1000.0}
with self.assertRaisesRegex(ValueError, "stale"):
require_fresh_selection(stale, max_age_seconds=60, now=1100.0, path="selection.json")
config = load_remote_run_config("configs/remote_smoke.toml")
manifest = choose_offer([offer(20, price=0.20, host=3)], config, query={}).to_manifest()
self.assertIn("created_at", manifest)
self.assertIn("created_at_iso", manifest)
self.assertIn("age_seconds", manifest)
class CleanupReconciliationTests(unittest.TestCase):
def test_reconciliation_uses_vast_ground_truth_for_orphans_and_records_actions(self) -> None:
destroyed: list[int] = []
report = reconcile_cleanup(
sky_state={"clusters": [{"name": "known-run", "instance_id": 77}]},
vast_instances=[
{"id": 77, "actual_status": "running", "gpu_name": "RTX 4090", "num_gpus": 1, "dph_total": 0.40},
{"id": 88, "host_id": 123, "actual_status": "running", "gpu_name": "RTX 4090", "num_gpus": 1, "dph_total": 0.45, "label": "orphan-run"},
],
known_run_ids=("known-run", "orphan-run"),
destroy_orphans=True,
destroy_instance=lambda instance_id: destroyed.append(instance_id),
now=1234.0,
)
orphan = next(item for item in report["instances"] if item["vast_instance_id"] == 88)
self.assertEqual(report["unexpected_live_count"], 1)
self.assertEqual(destroyed, [88])
self.assertEqual(orphan["cleanup_action_attempted"], "destroy_orphan")
self.assertEqual(orphan["cleanup_result"], "destroy_requested")
self.assertEqual(orphan["hourly_cost"], 0.45)
class HfSafetyTests(unittest.TestCase):
def test_rate_limit_suppression_preserves_training_success_as_hf_incomplete(self) -> None:
class FakeRateLimitError(RuntimeError):
def __init__(self) -> None:
super().__init__("429 Too Many Requests")
self.response = types.SimpleNamespace(headers={"Retry-After": "600"})
class FakeCommitOperationAdd:
def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None:
self.path_in_repo = path_in_repo
self.path_or_fileobj = path_or_fileobj
fake_module = types.SimpleNamespace(CommitOperationAdd=FakeCommitOperationAdd)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}):
run_dir = Path(tmp)
(run_dir / "metrics.jsonl").write_text("{}\n")
uploader = HfArtifactUploader(
enabled=True,
run_dir=run_dir,
repo_id="owner/repo",
repo_type="model",
path_in_repo="runs/run-1",
max_rate_limit_sleep_seconds=0,
)
uploader._api = types.SimpleNamespace(create_commit=Mock(side_effect=FakeRateLimitError()))
with self.assertRaises(FakeRateLimitError):
uploader.upload_files(("metrics.jsonl",), commit_message="upload metrics")
suppressed = uploader.upload_files(("metrics.jsonl",), commit_message="retry metrics")
final = uploader.finalize(training_success=True)
self.assertTrue(suppressed["rate_limited"])
self.assertEqual(final["hf_publication_status"], "training_succeeded_hf_incomplete")
manifest = json.loads((run_dir / "hf_upload_manifest.json").read_text())
self.assertTrue(manifest["training_success"])
self.assertFalse(manifest["publication_complete"])
self.assertGreater(manifest["rate_limit_until"], 0)
def test_final_reporting_distinguishes_training_failure_from_hf_success(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
disabled = HfArtifactUploader(enabled=False, run_dir=Path(tmp))
self.assertEqual(disabled.finalize(training_success=False)["hf_publication_status"], "disabled")
with tempfile.TemporaryDirectory() as tmp:
uploader = HfArtifactUploader(enabled=True, run_dir=Path(tmp), repo_id="owner/repo", repo_type="model", path_in_repo="run")
self.assertEqual(uploader.finalize(training_success=False)["hf_publication_status"], "training_failed")
with tempfile.TemporaryDirectory() as tmp:
uploader = HfArtifactUploader(enabled=True, run_dir=Path(tmp), repo_id="owner/repo", repo_type="model", path_in_repo="run")
self.assertEqual(uploader.finalize(training_success=True)["hf_publication_status"], "hf_publication_succeeded")
class ArtifactCollectionReportTests(unittest.TestCase):
def test_collection_report_classifies_produced_missing_partial_and_failed_copy(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
def copy_one(relative_path: str) -> int | None:
if relative_path == "produced.json":
(root / relative_path).write_text("{}\n")
return 0
if relative_path == "missing.json":
return 0
if relative_path == "partial.pt":
partial = root / ".rsync-partial" / relative_path
partial.parent.mkdir(parents=True)
partial.write_bytes(b"partial")
return 0
raise RuntimeError("rsync failed")
report = collect_artifact_paths(
local_dir=root,
remote_dir="remote:~/artifacts",
paths=("produced.json", "missing.json", "partial.pt", "failed.json"),
required=("produced.json", "partial.pt", "failed.json"),
collection_kind="terminal",
copy_one=copy_one,
)
by_path = {attempt["expected_path"]: attempt for attempt in report["attempts"]}
self.assertEqual(by_path["produced.json"]["final_status"], "success")
self.assertEqual(by_path["missing.json"]["likely_reason"], "remote_missing_or_not_produced")
self.assertEqual(by_path["partial.pt"]["final_status"], "partial")
self.assertEqual(by_path["failed.json"]["likely_reason"], "collection_command_failed")
self.assertEqual(
{item["expected_path"] for item in required_collection_failures(report)},
{"partial.pt", "failed.json"},
)
saved = json.loads((root / ARTIFACT_COLLECTION_REPORT).read_text())
self.assertFalse(saved["summary"]["ok"])
class DiskPhilosophyTests(unittest.TestCase):
def test_disk_paths_record_telemetry_and_backpressure_state_instead_of_capacity_mismatch_hard_fail(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml")
yaml = render_skypilot_yaml(config, choose_offer([offer(30, price=0.20, host=4)], config, query={}), run_id="disk-check")
self.assertIn("disk_telemetry.json", yaml)
self.assertIn("backpressure_adaptive", yaml)
self.assertIn("airfrans_disk_capacity_status=below_requested", yaml)
self.assertNotIn("exit 74", yaml)
with tempfile.TemporaryDirectory() as tmp:
run_dir = Path(tmp)
recorder = StreamingEventRecorder(run_dir)
usage = shutil._ntuple_diskusage(total=1000, used=900, free=100)
with patch("airfrans_frontier.training.streaming_data.shutil.disk_usage", return_value=usage):
recorder.observe_cache(run_dir, cache_bytes=950)
recorder.emit("cache_high_water", phase="data", cache_bytes=950, high_water_bytes=900)
recorder.emit("producer_paused", phase="data", reason="cache_high_water")
recorder.emit("cache_low_water", phase="data", cache_bytes=500, low_water_bytes=600)
recorder.emit("producer_resumed", phase="data", reason="cache_low_water", idle_seconds=1.25)
summary = recorder.to_dict()
self.assertEqual(summary["minimum_free_disk_bytes"], 100)
self.assertEqual(summary["cache_high_water_events"], 1)
self.assertEqual(summary["cache_low_water_events"], 1)
self.assertEqual(summary["producer_pause_events"], 1)
self.assertEqual(summary["producer_resume_events"], 1)
self.assertGreater(summary["producer_idle_backpressure_seconds"], 0)
def offer(offer_id: int, *, price: float, host: int) -> VastOffer:
return VastOffer(
id=offer_id,
gpu_name="RTX 4090",
dph_total=price,
gpu_ram=24_000,
disk_space=256.0,
geolocation="US",
inet_down_cost_per_tb=0.0,
inet_up_cost_per_tb=0.0,
host_id=host,
verification="verified",
reliability2=0.99,
cuda_max_good=12.8,
direct_port_count=1,
inet_down=500.0,
inet_up=100.0,
verified=True,
)
if __name__ == "__main__":
unittest.main()

View file

@ -1,161 +0,0 @@
from __future__ import annotations
import gzip
import sys
import tempfile
import types
import shutil
import unittest
import zipfile
from pathlib import Path
from unittest.mock import patch
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
from airfrans_frontier.raw.public import ensure_public_airfrans_processed_hf, extract_of_dataset, process_of_dataset_url_streaming
def write_minimal_airfrans_archive(archive: Path, case_names: list[str]) -> None:
with zipfile.ZipFile(archive, "w") as zf:
for case_name in case_names:
base = f"OF_dataset/{case_name}"
zf.writestr(f"{base}/constant/transportProperties", "nu 1e-5;\n")
zf.writestr(
f"{base}/constant/polyMesh/boundary",
"\naerofoil\n{\n type wall;\n nFaces 1;\n startFace 0;\n}\nfarfield\n{\n type patch;\n nFaces 3;\n startFace 1;\n}\n",
)
zf.writestr(f"{base}/constant/polyMesh/points.gz", gzip.compress(b"4\n(\n(0 0 0)\n(1 0 0)\n(1 1 0)\n(0 1 0)\n)\n"))
zf.writestr(f"{base}/constant/polyMesh/faces.gz", gzip.compress(b"4\n(\n2(0 1)\n2(1 2)\n2(2 3)\n2(3 0)\n)\n"))
zf.writestr(f"{base}/constant/polyMesh/owner.gz", gzip.compress(b"4\n(\n0\n0\n0\n0\n)\n"))
zf.writestr(f"{base}/constant/polyMesh/neighbour.gz", gzip.compress(b"0\n(\n)\n"))
zf.writestr(f"{base}/1/U.gz", gzip.compress(b"1\n(\n(1 0 0)\n)\n"))
zf.writestr(f"{base}/1/p.gz", gzip.compress(b"1\n(\n0.5\n)\n"))
zf.writestr(f"{base}/1/nut.gz", gzip.compress(b"1\n(\n0.01\n)\n"))
class PublicAirfransDataTests(unittest.TestCase):
def test_prepare_public_hf_skips_when_dataset_already_published(self) -> None:
class FakeApi:
def __init__(self, token=None):
self.token = token
def list_repo_files(self, *, repo_id: str, repo_type: str):
assert repo_id == "owner/airfrans-processed"
assert repo_type == "dataset"
return [
"processed/full/case_000.npz",
"processed/full/case_001.npz",
"processed/full/hf_dataset_manifest.json",
]
fake_module = types.SimpleNamespace(HfApi=FakeApi)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(
"os.environ", {"HF_TOKEN": "token"}
):
report = ensure_public_airfrans_processed_hf(
repo_id="owner/airfrans-processed",
path_in_repo="processed/full",
work_dir=Path(tmp) / "work",
output_dir=Path(tmp) / "out",
min_cases=2,
)
self.assertTrue(report["ok"])
self.assertEqual(report["phase"], "already_published")
self.assertEqual(report["npz_file_count"], 2)
self.assertTrue(report["has_manifest"])
def test_extract_of_dataset_finds_public_archive_root(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr("OF_dataset/airFoil2D_SST_demo/system/controlDict", "ok")
root = extract_of_dataset(archive, tmp_path / "raw", min_cases=1)
self.assertEqual(root.name, "OF_dataset")
def test_extract_of_dataset_rejects_zip_slip_paths(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "bad.zip"
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr("../escape.txt", "bad")
with self.assertRaisesRegex(RuntimeError, "Unsafe path"):
extract_of_dataset(archive, tmp_path / "raw", min_cases=1)
def test_extract_of_dataset_fails_before_partial_extract_when_disk_is_too_small(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr("OF_dataset/airFoil2D_SST_demo/system/controlDict", "ok")
tiny_disk = shutil._ntuple_diskusage(total=10, used=10, free=0)
with patch("airfrans_frontier.raw.public.shutil.disk_usage", return_value=tiny_disk):
with self.assertRaisesRegex(RuntimeError, "Insufficient free disk"):
extract_of_dataset(archive, tmp_path / "raw", min_cases=1)
self.assertFalse((tmp_path / "raw" / "OF_dataset").exists())
def test_range_streaming_processing_writes_npz_and_discards_raw_case(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
case_name = "airFoil2D_SST_10.0_5.0_0012"
write_minimal_airfrans_archive(archive, [case_name])
streamed = process_of_dataset_url_streaming(
str(archive),
tmp_path / "processed",
scratch_dir=tmp_path / "streaming_raw",
min_cases=1,
progress_every=1,
)
result = streamed.processing
self.assertEqual(result.case_count, 1)
self.assertTrue((tmp_path / "processed" / f"{case_name}.npz").is_file())
self.assertTrue(result.manifest_path.is_file())
self.assertFalse((tmp_path / "streaming_raw" / case_name).exists())
self.assertGreater(streamed.ranged_bytes_read, 0)
def test_prepare_public_hf_streams_archive_before_publish(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
source_archive = tmp_path / "source_OF_dataset.zip"
case_name = "airFoil2D_SST_10.0_5.0_0012"
write_minimal_airfrans_archive(source_archive, [case_name])
statuses = [
{"file_count": 0, "npz_file_count": 0, "has_manifest": False},
{"file_count": 2, "npz_file_count": 1, "has_manifest": True},
]
def fake_publish(**kwargs):
data_root = Path(kwargs["data_root"])
self.assertTrue((data_root / f"{case_name}.npz").is_file())
self.assertFalse((tmp_path / "work" / "streaming_raw" / case_name).exists())
return {"repo_url": "https://huggingface.co/datasets/owner/repo", "npz_file_count": 1}
with patch("airfrans_frontier.raw.public._hf_dataset_status", side_effect=statuses), patch(
"airfrans_frontier.raw.public.publish_processed_dataset", side_effect=fake_publish
):
report = ensure_public_airfrans_processed_hf(
repo_id="owner/repo",
path_in_repo="processed/full",
work_dir=tmp_path / "work",
output_dir=tmp_path / "processed",
source_url=str(source_archive),
min_cases=1,
)
self.assertFalse((tmp_path / "work" / "OF_dataset.zip").exists())
self.assertTrue(report["ok"])
self.assertTrue(report["streaming"])
self.assertEqual(report["streaming_mode"], "zip_range")
self.assertEqual(report["download"]["mode"], "zip_range")
self.assertEqual(report["processed_case_count"], 1)
if __name__ == "__main__":
unittest.main()

View file

@ -1,22 +1,15 @@
from __future__ import annotations from __future__ import annotations
from contextlib import redirect_stdout
from io import StringIO
import json import json
import tempfile import tempfile
import shutil import shutil
import unittest import unittest
from unittest.mock import patch
from pathlib import Path from pathlib import Path
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
import torch import torch
from airfrans_frontier.remote.artifacts import verify_artifacts from airfrans_frontier.remote.artifacts import verify_artifacts
from airfrans_frontier.remote.cli import _classify_artifacts, _stage_resume_checkpoint, _terminal_artifact_names, main as remote_main from airfrans_frontier.remote.cli import _classify_artifacts, _stage_resume_checkpoint
from airfrans_frontier.remote.config import load_remote_run_config from airfrans_frontier.remote.config import load_remote_run_config
from airfrans_frontier.remote.skypilot import render_skypilot_yaml from airfrans_frontier.remote.skypilot import render_skypilot_yaml
from airfrans_frontier.remote.vast import VastOffer, choose_offer from airfrans_frontier.remote.vast import VastOffer, choose_offer
@ -82,16 +75,6 @@ class VastSelectionTests(unittest.TestCase):
self.assertNotIn("sky launch", yaml) self.assertNotIn("sky launch", yaml)
self.assertIn("remote-run smoke-train", yaml) self.assertIn("remote-run smoke-train", yaml)
self.assertIn("configs/aggressive_smoke.toml", yaml) self.assertIn("configs/aggressive_smoke.toml", yaml)
self.assertIn("df -h .", yaml)
self.assertIn("airfrans_disk_requested_gb=128", yaml)
self.assertIn("AIRFRANS_STARTUP_TIMELINE: artifacts/current_run/startup_timeline.jsonl", yaml)
self.assertIn("airfrans_timeline 'setup' 'started'", yaml)
self.assertIn("airfrans_timeline 'data_validation' 'started'", yaml)
self.assertIn("airfrans_timeline 'training_command' 'started'", yaml)
def test_terminal_collection_includes_startup_timeline(self) -> None:
self.assertIn("startup_timeline.jsonl", _terminal_artifact_names(()))
def test_rendered_yaml_can_pass_resume_checkpoint(self) -> None: def test_rendered_yaml_can_pass_resume_checkpoint(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml") config = load_remote_run_config("configs/remote_smoke.toml")
@ -107,22 +90,6 @@ class VastSelectionTests(unittest.TestCase):
self.assertIn("AIRFRANS_RESUME_CHECKPOINT: .airfrans_resume/airfrans-test/checkpoint_latest.pt", yaml) self.assertIn("AIRFRANS_RESUME_CHECKPOINT: .airfrans_resume/airfrans-test/checkpoint_latest.pt", yaml)
class VastInstanceCliTests(unittest.TestCase):
def test_vast_instances_reports_api_ground_truth(self) -> None:
stdout = StringIO()
with patch.dict("os.environ", {"VAST_API_KEY": "token"}), patch(
"airfrans_frontier.remote.cli.list_instances",
return_value=[{"id": 123, "actual_status": "running", "gpu_name": "RTX 4090"}],
), redirect_stdout(stdout):
code = remote_main(["vast-instances"])
self.assertEqual(code, 0)
payload = json.loads(stdout.getvalue())
self.assertEqual(payload["instance_count"], 1)
self.assertEqual(payload["instances"][0]["id"], 123)
self.assertEqual(payload["instances"][0]["actual_status"], "running")
class ArtifactVerificationTests(unittest.TestCase): class ArtifactVerificationTests(unittest.TestCase):
def test_verify_artifacts_requires_contract_files_and_writes_manifest(self) -> None: def test_verify_artifacts_requires_contract_files_and_writes_manifest(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
@ -134,9 +101,6 @@ class ArtifactVerificationTests(unittest.TestCase):
self.assertGreaterEqual(manifest["file_count"], 8) self.assertGreaterEqual(manifest["file_count"], 8)
self.assertTrue((root / "artifact_manifest.json").is_file()) self.assertTrue((root / "artifact_manifest.json").is_file())
self.assertTrue((root / "checksums.txt").is_file()) self.assertTrue((root / "checksums.txt").is_file())
self.assertTrue((root / "verification_report.json").is_file())
report = json.loads((root / "verification_report.json").read_text())
self.assertTrue(report["ok"])
def test_verify_artifacts_accepts_failure_report_terminal_state(self) -> None: def test_verify_artifacts_accepts_failure_report_terminal_state(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
@ -203,7 +167,6 @@ def offer(
gpu_name="RTX 4090", gpu_name="RTX 4090",
dph_total=price, dph_total=price,
gpu_ram=24_000, gpu_ram=24_000,
disk_space=256.0,
geolocation=geo, geolocation=geo,
inet_down_cost_per_tb=0.0, inet_down_cost_per_tb=0.0,
inet_up_cost_per_tb=0.0, inet_up_cost_per_tb=0.0,

View file

@ -1,51 +0,0 @@
from __future__ import annotations
import json
import sys
import tempfile
import types
import unittest
from pathlib import Path
from unittest.mock import patch
from airfrans_frontier.remote.smoke import run_smoke_training
class SmokeTrainingFailureTests(unittest.TestCase):
def test_pre_checkpoint_training_error_writes_terminal_failure_report(self) -> None:
fake_loop = types.ModuleType("airfrans_frontier.training.loop")
def fail_train(*_: object, **__: object) -> object:
raise RuntimeError("hub commit rate limited")
fake_loop.train_from_config_path = fail_train
fake_torch = types.ModuleType("torch")
fake_torch.__version__ = "fake"
fake_torch.version = types.SimpleNamespace(cuda=None)
fake_torch.cuda = types.SimpleNamespace(
is_available=lambda: False,
get_device_name=lambda _index: None,
)
with tempfile.TemporaryDirectory() as tmp, patch.dict(
sys.modules,
{
"airfrans_frontier.training.loop": fake_loop,
"torch": fake_torch,
},
):
artifact_dir = Path(tmp)
with self.assertRaisesRegex(RuntimeError, "hub commit rate limited"):
run_smoke_training("missing-config.toml", artifact_dir=artifact_dir, run_id="smoke-fail")
report = json.loads((artifact_dir / "failure_report.json").read_text())
self.assertEqual(report["run_id"], "smoke-fail")
self.assertEqual(report["error_type"], "RuntimeError")
self.assertEqual(report["error_message"], "hub commit rate limited")
verification = json.loads((artifact_dir / "verification_report.json").read_text())
self.assertTrue(verification["ok"])
self.assertEqual(verification["checks"]["terminal_artifact"], "failure_report.json")
if __name__ == "__main__":
unittest.main()

View file

@ -1,372 +0,0 @@
from __future__ import annotations
import gzip
import json
import os
import sys
import tempfile
import types
import unittest
import zipfile
from pathlib import Path
from unittest.mock import patch
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
import numpy as np
from airfrans_frontier.raw.public import process_of_dataset_url_streaming
from airfrans_frontier.training.config import load_training_config
from airfrans_frontier.training.data import build_dataset_bundle, load_processed_dataset
from airfrans_frontier.training.loop import train
from airfrans_frontier.training.normalize import compute_normalization_stats
from airfrans_frontier.training.streaming_data import StreamingEventRecorder, StreamingTrainingData
def write_minimal_airfrans_archive(archive: Path, case_names: list[str]) -> None:
with zipfile.ZipFile(archive, "w") as zf:
for index, case_name in enumerate(case_names):
base = f"OF_dataset/{case_name}"
u_value = 1.0 + 0.1 * index
p_value = 0.5 + 0.2 * index
nut_value = 0.01 + 0.001 * index
zf.writestr(f"{base}/constant/transportProperties", "nu 1e-5;\n")
zf.writestr(
f"{base}/constant/polyMesh/boundary",
"\naerofoil\n{\n type wall;\n nFaces 1;\n startFace 0;\n}\nfarfield\n{\n type patch;\n nFaces 3;\n startFace 1;\n}\n",
)
zf.writestr(f"{base}/constant/polyMesh/points.gz", gzip.compress(b"4\n(\n(0 0 0)\n(1 0 0)\n(1 1 0)\n(0 1 0)\n)\n"))
zf.writestr(f"{base}/constant/polyMesh/faces.gz", gzip.compress(b"4\n(\n2(0 1)\n2(1 2)\n2(2 3)\n2(3 0)\n)\n"))
zf.writestr(f"{base}/constant/polyMesh/owner.gz", gzip.compress(b"4\n(\n0\n0\n0\n0\n)\n"))
zf.writestr(f"{base}/constant/polyMesh/neighbour.gz", gzip.compress(b"0\n(\n)\n"))
zf.writestr(f"{base}/1/U.gz", gzip.compress(f"1\n(\n({u_value} 0 0)\n)\n".encode()))
zf.writestr(f"{base}/1/p.gz", gzip.compress(f"1\n(\n{p_value}\n)\n".encode()))
zf.writestr(f"{base}/1/nut.gz", gzip.compress(f"1\n(\n{nut_value}\n)\n".encode()))
def write_malformed_airfrans_archive(archive: Path, case_name: str) -> None:
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr(f"OF_dataset/{case_name}/constant/transportProperties", "nu 1e-5;\n")
def write_streaming_config(
path: Path,
*,
archive: Path,
cache_dir: Path,
artifact_dir: Path,
train_cases: int = 2,
val_cases: int = 1,
test_cases: int = 1,
steps: int = 2,
log_interval: int = 1,
batch_size: int = 2,
high_water_bytes: int = 32 * 1024 * 1024,
low_water_bytes: int = 16 * 1024 * 1024,
upload_processed: bool = False,
upload_batch_size: int = 1,
) -> None:
path.write_text(
f"""
[run]
name = "streaming_test"
seed = 7
artifact_dir = "{artifact_dir}"
[data]
root = "{cache_dir}"
source = "public_zip_streaming"
public_source_url = "{archive}"
cache_dir = "{cache_dir}"
streaming_scratch_dir = "{cache_dir / '_raw'}"
train_cases = {train_cases}
val_cases = {val_cases}
test_cases = {test_cases}
points_per_case = 999999999
batch_size = {batch_size}
streaming_cache_max_bytes = {max(high_water_bytes, high_water_bytes + 1)}
streaming_cache_high_water_bytes = {high_water_bytes}
streaming_cache_low_water_bytes = {low_water_bytes}
streaming_queue_max_cases = 1
streaming_upload_processed = {str(upload_processed).lower()}
streaming_upload_batch_size = {upload_batch_size}
hf_repo_id = "owner/airfrans-processed"
hf_repo_type = "dataset"
hf_path_prefix = "processed/full"
[model]
type = "mlp"
hidden_width = 16
depth = 2
activation = "gelu"
[optim]
lr = 0.01
weight_decay = 0.0
steps = {steps}
log_interval = {log_interval}
[device]
type = "cpu"
allow_cpu_fallback = false
benchmark_kernels = false
[loss]
type = "normalized_mse"
[checkpoint]
interval_seconds = 0
""".strip()
+ "\n"
)
def read_events(run_dir: Path) -> list[dict[str, object]]:
return [json.loads(line) for line in (run_dir / "streaming_events.jsonl").read_text().splitlines() if line.strip()]
class FullDataBackpressureStreamingTests(unittest.TestCase):
def test_streaming_training_smoke_writes_artifacts_without_eager_concatenation(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(5)]
write_minimal_airfrans_archive(archive, case_names)
config_path = tmp_path / "streaming.toml"
artifact_dir = tmp_path / "artifacts"
write_streaming_config(config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=artifact_dir)
config = load_training_config(config_path)
with patch("airfrans_frontier.training.loop.load_processed_dataset", side_effect=AssertionError("eager load called")), patch(
"airfrans_frontier.training.loop.build_dataset_bundle", side_effect=AssertionError("eager concat called")
):
result = train(config)
self.assertTrue(np.isfinite(result.final_metrics["train_loss"]))
self.assertEqual(result.final_metrics["data_mode"], "public_zip_streaming")
for name in (
"metrics.jsonl",
"checkpoint_latest.pt",
"checkpoint_best.pt",
"checkpoint_final.pt",
"final_metrics.json",
"split_manifest.json",
"data_manifest.json",
"normalization.json",
"streaming_events.jsonl",
"streaming_state.json",
"streaming_summary.json",
"processed_upload_manifest.json",
"artifact_manifest.json",
"checksums.txt",
"verification_report.json",
):
self.assertTrue((result.run_dir / name).is_file(), name)
events = read_events(result.run_dir)
event_names = {event["event"] for event in events}
self.assertIn("dataset_enumeration_start", event_names)
self.assertIn("dataset_enumeration_end", event_names)
self.assertIn("split_selection", event_names)
self.assertIn("normalization_start", event_names)
self.assertIn("normalization_end", event_names)
self.assertIn("first_batch_ready", event_names)
self.assertIn("first_gpu_batch_consumed", event_names)
self.assertIn("first_metric", event_names)
self.assertIn("first_checkpoint_written", event_names)
selected_cases = set(json.loads((result.run_dir / "data_manifest.json").read_text())["cases"][index]["case_id"] for index in range(4))
processed_cases = {str(event["case_id"]) for event in events if event["event"] == "processing_end"}
self.assertLessEqual(processed_cases, selected_cases)
self.assertFalse(any((tmp_path / "cache" / "_raw").glob("airFoil2D_*")))
def test_backpressure_pauses_resumes_and_bounds_cache_with_inflight_slack(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(4)]
write_minimal_airfrans_archive(archive, case_names)
config_path = tmp_path / "streaming.toml"
high_water = 256
write_streaming_config(
config_path,
archive=archive,
cache_dir=tmp_path / "cache",
artifact_dir=tmp_path / "artifacts",
high_water_bytes=high_water,
low_water_bytes=128,
steps=2,
)
result = train(load_training_config(config_path))
summary = json.loads((result.run_dir / "streaming_summary.json").read_text())
self.assertGreater(summary["cache_high_water_events"], 0)
self.assertGreater(summary["cache_low_water_events"], 0)
self.assertGreater(summary["producer_pause_events"], 0)
self.assertGreater(summary["producer_resume_events"], 0)
self.assertGreater(summary["evicted_units"], 0)
self.assertLessEqual(summary["processed_cache_high_water_bytes"], high_water + summary["max_processed_unit_bytes"])
event_names = {event["event"] for event in read_events(result.run_dir)}
self.assertIn("producer_paused", event_names)
self.assertIn("producer_resumed", event_names)
self.assertIn("cleanup_eviction", event_names)
def test_streaming_normalization_matches_eager_train_split_statistics(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(4)]
write_minimal_airfrans_archive(archive, case_names)
config_path = tmp_path / "streaming.toml"
write_streaming_config(config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=tmp_path / "artifacts", steps=1)
config = load_training_config(config_path)
run_dir = tmp_path / "run"
recorder = StreamingEventRecorder(run_dir)
streaming = StreamingTrainingData.from_config(config, run_dir=run_dir, recorder=recorder)
streaming.prepare()
streaming_stats = streaming.load_or_compute_normalization()
eager_root = tmp_path / "eager_processed"
process_of_dataset_url_streaming(str(archive), eager_root, scratch_dir=tmp_path / "eager_raw", min_cases=4)
eager_bundle = build_dataset_bundle(
load_processed_dataset(eager_root),
train_cases=config.data.train_cases,
val_cases=config.data.val_cases,
test_cases=config.data.test_cases,
points_per_case=config.data.points_per_case,
seed=config.run.seed,
)
eager_stats = compute_normalization_stats(
eager_bundle.train.features,
eager_bundle.train.targets,
feature_names=eager_bundle.feature_names,
target_names=eager_bundle.target_names,
)
np.testing.assert_allclose(streaming_stats.feature_mean, eager_stats.feature_mean, rtol=1e-6, atol=1e-6)
np.testing.assert_allclose(streaming_stats.feature_std, eager_stats.feature_std, rtol=1e-6, atol=1e-6)
np.testing.assert_allclose(streaming_stats.target_mean, eager_stats.target_mean, rtol=1e-6, atol=1e-6)
np.testing.assert_allclose(streaming_stats.target_std, eager_stats.target_std, rtol=1e-6, atol=1e-6)
def test_resume_reuses_validated_units_and_discards_partial_units(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(3)]
write_minimal_airfrans_archive(archive, case_names)
config_path = tmp_path / "streaming.toml"
write_streaming_config(
config_path,
archive=archive,
cache_dir=tmp_path / "cache",
artifact_dir=tmp_path / "artifacts",
train_cases=1,
val_cases=1,
test_cases=1,
steps=1,
)
config = load_training_config(config_path)
run_dir = tmp_path / "run"
first = StreamingTrainingData.from_config(config, run_dir=run_dir, recorder=StreamingEventRecorder(run_dir))
first.prepare()
assert first.split is not None
first_case = first.split.train_ids[0]
(tmp_path / "cache" / f"{first_case}.npz.tmp.npz").write_bytes(b"partial")
second = StreamingTrainingData.from_config(config, run_dir=run_dir, recorder=StreamingEventRecorder(run_dir))
second.prepare()
events = read_events(run_dir)
self.assertTrue(any(event["event"] == "partial_unit_discarded" and event.get("case_id") == first_case for event in events))
self.assertTrue(any(event["event"] == "resume_validated_unit_reused" and event.get("case_id") == first_case for event in events))
processing_events = [event for event in events if event["event"] == "processing_end" and event.get("case_id") == first_case]
self.assertEqual(len(processing_events), 1)
def test_processed_upload_rate_limit_does_not_fail_training(self) -> None:
class FakeRateLimitError(RuntimeError):
def __init__(self) -> None:
super().__init__("429 Too Many Requests")
self.response = types.SimpleNamespace(headers={"Retry-After": "600"})
class FakeCommitOperationAdd:
def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None:
self.path_in_repo = path_in_repo
self.path_or_fileobj = path_or_fileobj
class FakeApi:
def __init__(self, token: str) -> None:
self.token = token
def create_repo(self, *, repo_id: str, repo_type: str, private: bool, exist_ok: bool) -> None:
return None
def create_commit(self, **kwargs):
raise FakeRateLimitError()
fake_module = types.SimpleNamespace(HfApi=FakeApi, CommitOperationAdd=FakeCommitOperationAdd)
with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}):
tmp_path = Path(tmp)
archive = tmp_path / "OF_dataset.zip"
case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(4)]
write_minimal_airfrans_archive(archive, case_names)
config_path = tmp_path / "streaming.toml"
write_streaming_config(
config_path,
archive=archive,
cache_dir=tmp_path / "cache",
artifact_dir=tmp_path / "artifacts",
upload_processed=True,
upload_batch_size=1,
steps=1,
)
result = train(load_training_config(config_path))
self.assertTrue(np.isfinite(result.final_metrics["train_loss"]))
manifest = json.loads((result.run_dir / "processed_upload_manifest.json").read_text())
self.assertTrue(manifest["enabled"])
self.assertGreater(manifest["queue_depth"], 0)
self.assertGreater(manifest["rate_limit_until"], 0)
self.assertEqual(manifest["rate_limit_retry_after_seconds"], 600.0)
events = {event["event"] for event in read_events(result.run_dir)}
self.assertIn("processed_data_upload_rate_limited", events)
self.assertIn("processed_data_upload_suppressed", events)
run_manifest = json.loads((result.run_dir / "run_manifest.json").read_text())
self.assertEqual(run_manifest["phase"], "completed")
def test_streaming_failure_writes_diagnostic_artifacts(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp)
archive = tmp_path / "bad.zip"
case_name = "airFoil2D_SST_10.0_5.0_0012"
write_malformed_airfrans_archive(archive, case_name)
config_path = tmp_path / "streaming.toml"
artifact_dir = tmp_path / "artifacts"
write_streaming_config(
config_path,
archive=archive,
cache_dir=tmp_path / "cache",
artifact_dir=artifact_dir,
train_cases=1,
val_cases=0,
test_cases=0,
steps=1,
)
with self.assertRaises(Exception):
train(load_training_config(config_path))
run_dir = next(path for path in artifact_dir.iterdir() if path.is_dir())
for name in ("failure_report.json", "metrics.jsonl", "streaming_events.jsonl", "streaming_state.json", "streaming_summary.json", "verification_report.json"):
self.assertTrue((run_dir / name).is_file(), name)
report = json.loads((run_dir / "failure_report.json").read_text())
self.assertEqual(report["phase"], "streaming_training")
events = {event["event"] for event in read_events(run_dir)}
self.assertIn("processing_failure", events)
verification = json.loads((run_dir / "verification_report.json").read_text())
self.assertFalse(verification["ok"])
if __name__ == "__main__":
unittest.main()

View file

@ -16,17 +16,6 @@ class TrainingConfigTests(unittest.TestCase):
self.assertEqual(config.loss.type, "normalized_mse") self.assertEqual(config.loss.type, "normalized_mse")
self.assertEqual(config.device.type, "cuda") self.assertEqual(config.device.type, "cuda")
self.assertTrue(config.data.root.is_absolute()) self.assertTrue(config.data.root.is_absolute())
self.assertEqual(config.data.source, "local")
self.assertIsNone(config.data.hf_repo_id)
self.assertIsNone(config.data.cache_dir)
def test_config_loader_accepts_huggingface_data_source(self) -> None:
config = load_training_config("configs/aggressive_smoke.toml")
self.assertEqual(config.data.source, "huggingface")
self.assertEqual(config.data.hf_repo_id, "zacheryasc/airfrans-processed")
self.assertEqual(config.data.hf_path_prefix, "processed/full")
self.assertTrue(config.data.cache_dir is not None)
def test_config_loader_rejects_missing_section(self) -> None: def test_config_loader_rejects_missing_section(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:

View file

@ -18,7 +18,6 @@ import torch
from unittest.mock import Mock from unittest.mock import Mock
from airfrans_frontier.training.config import load_training_config from airfrans_frontier.training.config import load_training_config
from airfrans_frontier.training.hf_upload import HfArtifactUploader
from airfrans_frontier.training.loop import train, select_device from airfrans_frontier.training.loop import train, select_device
@ -108,40 +107,6 @@ interval_seconds = {checkpoint_interval_seconds}
) )
class HfArtifactUploaderTests(unittest.TestCase):
def test_upload_files_commits_batch_once_to_reduce_hub_rate_limit(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
run_dir = Path(tmp)
(run_dir / "metrics.jsonl").write_text("{}\n")
(run_dir / "latest_metrics.json").write_text("{}\n")
uploader = HfArtifactUploader(
enabled=True,
run_dir=run_dir,
repo_id="owner/repo",
repo_type="model",
path_in_repo="runs/model",
)
fake_api = Mock()
fake_api.create_commit.return_value = Mock(oid="abc123", commit_url="https://hf/commit/abc123", pr_url=None)
uploader._api = fake_api
result = uploader.upload_files(("metrics.jsonl", "latest_metrics.json"), commit_message="batch artifacts")
self.assertEqual(result["uploaded"], ["runs/model/metrics.jsonl", "runs/model/latest_metrics.json"])
fake_api.create_commit.assert_called_once()
call_kwargs = fake_api.create_commit.call_args.kwargs
self.assertEqual(call_kwargs["repo_id"], "owner/repo")
self.assertEqual(call_kwargs["repo_type"], "model")
self.assertEqual(call_kwargs["commit_message"], "batch artifacts")
self.assertEqual(
[operation.path_in_repo for operation in call_kwargs["operations"]],
["runs/model/metrics.jsonl", "runs/model/latest_metrics.json"],
)
manifest = json.loads((run_dir / "hf_upload_manifest.json").read_text())
self.assertEqual(len(manifest["commits"]), 1)
self.assertEqual(set(manifest["uploaded_paths"]), set(result["uploaded"]))
class TrainingLoopTests(unittest.TestCase): class TrainingLoopTests(unittest.TestCase):
def test_cuda_config_fails_clearly_when_cuda_unavailable(self) -> None: def test_cuda_config_fails_clearly_when_cuda_unavailable(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
@ -211,18 +176,11 @@ class TrainingLoopTests(unittest.TestCase):
"rng_state", "rng_state",
"torch_rng_state", "torch_rng_state",
"batch_rng_state", "batch_rng_state",
"scheduler_state_dict",
): ):
self.assertIn(key, checkpoint) self.assertIn(key, checkpoint)
self.assertEqual(list(run_dir.glob("*.tmp")), []) self.assertEqual(list(run_dir.glob("*.tmp")), [])
self.assertTrue((run_dir / "artifact_manifest.json").is_file()) self.assertTrue((run_dir / "artifact_manifest.json").is_file())
self.assertTrue((run_dir / "checksums.txt").is_file()) self.assertTrue((run_dir / "checksums.txt").is_file())
self.assertTrue((run_dir / "verification_report.json").is_file())
self.assertTrue((run_dir / "hf_upload_manifest.json").is_file())
self.assertTrue((run_dir / "calibration_manifest.json").is_file())
self.assertTrue((run_dir / "environment_manifest.json").is_file())
self.assertTrue((run_dir / "evaluation_protocol.json").is_file())
self.assertTrue((run_dir / "run_manifest.json").is_file())
self.assertTrue((run_dir / "metrics.jsonl").exists()) self.assertTrue((run_dir / "metrics.jsonl").exists())
heartbeat = json.loads((live_dir / "heartbeat.json").read_text()) heartbeat = json.loads((live_dir / "heartbeat.json").read_text())
self.assertEqual(heartbeat["run_id"], "test-run") self.assertEqual(heartbeat["run_id"], "test-run")
@ -230,10 +188,6 @@ class TrainingLoopTests(unittest.TestCase):
self.assertTrue(np.isfinite(heartbeat["latest_metrics"]["train_loss"])) self.assertTrue(np.isfinite(heartbeat["latest_metrics"]["train_loss"]))
self.assertTrue((live_dir / "latest_metrics.json").is_file()) self.assertTrue((live_dir / "latest_metrics.json").is_file())
self.assertEqual(final_metrics["device"].startswith("cuda"), torch.cuda.is_available()) self.assertEqual(final_metrics["device"].startswith("cuda"), torch.cuda.is_available())
self.assertIn("estimated_forward_flops_per_item", final_metrics)
self.assertIn("estimated_train_flops", final_metrics)
self.assertIn("checkpoint_final_bytes", final_metrics)
self.assertFalse(final_metrics["context_target_values_allowed"])
if torch.cuda.is_available(): if torch.cuda.is_available():
self.assertIn("T550", final_metrics["gpu_name"]) self.assertIn("T550", final_metrics["gpu_name"])

349
uv.lock
View file

@ -199,7 +199,7 @@ dev = [
requires-dist = [ requires-dist = [
{ name = "huggingface-hub", specifier = ">=0.36.0" }, { name = "huggingface-hub", specifier = ">=0.36.0" },
{ name = "numpy", specifier = ">=2.4.0" }, { name = "numpy", specifier = ">=2.4.0" },
{ name = "torch", specifier = ">=2.7.1,<2.8.0" }, { name = "torch", specifier = ">=2.8.0" },
{ name = "wandb", specifier = ">=0.23.0" }, { name = "wandb", specifier = ">=0.23.0" },
] ]
@ -761,6 +761,83 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/aa/50/a9caea39ad19c431c1a3f8a31114df65b260cdfe67786b6c7e7c040c4c44/cryptography-49.0.0-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:be9fcb48a55f023493482827d4f459bd263cc20efde64f204b97c123201850c6", size = 3783731, upload-time = "2026-06-12T20:02:43.319Z" }, { url = "https://files.pythonhosted.org/packages/aa/50/a9caea39ad19c431c1a3f8a31114df65b260cdfe67786b6c7e7c040c4c44/cryptography-49.0.0-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:be9fcb48a55f023493482827d4f459bd263cc20efde64f204b97c123201850c6", size = 3783731, upload-time = "2026-06-12T20:02:43.319Z" },
] ]
[[package]]
name = "cuda-bindings"
version = "13.3.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cuda-pathfinder", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/51/6b/457ca12dad3ee9bfcc9a545cfd6b64b359ba49de40f776f6e028e678f262/cuda_bindings-13.3.1-cp311-cp311-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c5879712accf6e14bb01aa5e67440eb84998b8d104b509cc7a6dc0b8f656a474", size = 6053539, upload-time = "2026-05-29T23:11:43.19Z" },
{ url = "https://files.pythonhosted.org/packages/95/7a/c5e3c34a409b148f5c0f5a4ea374158f95d488862c1dffedf9aa5c639df9/cuda_bindings-13.3.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:04436a9364059c84b8f9636f359eccda1cf814341f5b670c71d80d2f79dbc708", size = 6674166, upload-time = "2026-05-29T23:11:45.478Z" },
{ url = "https://files.pythonhosted.org/packages/ce/67/5e7dba1ba576dd73da5dee894ca076ca5e959450dfff66d6d510a255d1f7/cuda_bindings-13.3.1-cp312-cp312-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c7855c4868aabc0cfae28abbe83d56734bdfbd08f08fc234ac1912a12858bf49", size = 6025351, upload-time = "2026-05-29T23:11:49.685Z" },
{ url = "https://files.pythonhosted.org/packages/39/2a/6d2e9047d1fb243dbaa364b01e0297534b9ed7fd27dba1c9f361519cf69b/cuda_bindings-13.3.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:e32d08f71ebcdf00f0f41eab2eb37e8da94c8ed411cc9f7f7a019ce6b34abe3a", size = 6657965, upload-time = "2026-05-29T23:11:52.227Z" },
{ url = "https://files.pythonhosted.org/packages/cc/6e/2394f8163360f8391f8f1b7e72d300a82724edb81a7b7084c799fbd4c91f/cuda_bindings-13.3.1-cp313-cp313-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9efb21c1ee64981e184b9e0ba5eb3179e5ba3d4b51665a6cb52b8ef3d01a7cbf", size = 5920504, upload-time = "2026-05-29T23:11:56.883Z" },
{ url = "https://files.pythonhosted.org/packages/34/c2/ef9b6a63f7dc432712a462c816662e662e00d38caa9b861c8c2588195d03/cuda_bindings-13.3.1-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2732904099e0a4d4db774a5fc6d91ee95fae065b4d2ecabb4968c5fe2406c9d7", size = 6476660, upload-time = "2026-05-29T23:11:59.188Z" },
{ url = "https://files.pythonhosted.org/packages/b1/81/bff68ce829999c1e4209c761bbf903b1c06ec570416ddb25020864ad5907/cuda_bindings-13.3.1-cp314-cp314-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:1ab2f74ed65bfef4163ba07a8db16f1085e0729291db12a2423aff84ee8278b8", size = 6013639, upload-time = "2026-05-29T23:12:03.509Z" },
{ url = "https://files.pythonhosted.org/packages/d4/e0/c8a1f0c8f9ffdea4f5fe6dbab89b326cef4d85caf489dad39e209da89416/cuda_bindings-13.3.1-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:efd4c814d311ec08c981f6dded1dbe7d4b371067ee4f6c14cccec4bde9590f80", size = 6534419, upload-time = "2026-05-29T23:12:05.633Z" },
{ url = "https://files.pythonhosted.org/packages/52/b8/83b1f563925b290f2d11a01a77a84013ba56052fe3653a5bef3ccfbb43d6/cuda_bindings-13.3.1-cp314-cp314t-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c3c772dfff49681541d59630c90f858e173ac926b9c593a2b7123f2a1043cc76", size = 5809771, upload-time = "2026-05-29T23:12:10.422Z" },
{ url = "https://files.pythonhosted.org/packages/12/20/e79b4bfe98f075195afb6343d41c498f9dbd2d161d7021d4d28bceb83581/cuda_bindings-13.3.1-cp314-cp314t-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:36febb7c1079d68a981dbbd8d5a67235b399802b82075c9388624719607e52b9", size = 6358584, upload-time = "2026-05-29T23:12:12.767Z" },
]
[[package]]
name = "cuda-pathfinder"
version = "1.5.6"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/d2/53/8fc9b0cdc5b7f62746e6a01b85b6461e5ae27f871010a5fcf8fa6950766d/cuda_pathfinder-1.5.6-py3-none-any.whl", hash = "sha256:7e4c07c117b78ba1fb35dac4c444d21f3677b1b1ff56175c53a8e3025c5b43c0", size = 52972, upload-time = "2026-06-30T00:58:04.34Z" },
]
[[package]]
name = "cuda-toolkit"
version = "13.0.3.0"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/d1/c7/a79086a62c98befcdb8349656c6f114e2db3b8b2422f6e25c97a7f2a9a3c/cuda_toolkit-13.0.3.0-py2.py3-none-any.whl", hash = "sha256:d693caaa261214ddd7dbb60d68e71cbed884e68c2be7509778f3051da0b91c3f", size = 2512, upload-time = "2026-04-14T00:50:08.173Z" },
]
[package.optional-dependencies]
cublas = [
{ name = "nvidia-cublas", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
{ name = "nvidia-cuda-nvrtc", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
cudart = [
{ name = "nvidia-cuda-runtime", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
cufft = [
{ name = "nvidia-cufft", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
{ name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
cufile = [
{ name = "nvidia-cufile", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
cupti = [
{ name = "nvidia-cuda-cupti", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
curand = [
{ name = "nvidia-curand", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
cusolver = [
{ name = "nvidia-cublas", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
{ name = "nvidia-cusolver", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
{ name = "nvidia-cusparse", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
{ name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
cusparse = [
{ name = "nvidia-cusparse", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
{ name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
nvjitlink = [
{ name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
nvrtc = [
{ name = "nvidia-cuda-nvrtc", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
nvtx = [
{ name = "nvidia-nvtx", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" },
]
[[package]] [[package]]
name = "cycler" name = "cycler"
version = "0.12.1" version = "0.12.1"
@ -2237,136 +2314,155 @@ wheels = [
] ]
[[package]] [[package]]
name = "nvidia-cublas-cu12" name = "nvidia-cublas"
version = "12.6.4.1" version = "13.1.1.3"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/af/eb/ff4b8c503fa1f1796679dce648854d58751982426e4e4b37d6fce49d259c/nvidia_cublas_cu12-12.6.4.1-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:08ed2686e9875d01b58e3cb379c6896df8e76c75e0d4a7f7dace3d7b6d9ef8eb", size = 393138322, upload-time = "2024-11-20T17:40:25.65Z" },
]
[[package]]
name = "nvidia-cuda-cupti-cu12"
version = "12.6.80"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/49/60/7b6497946d74bcf1de852a21824d63baad12cd417db4195fc1bfe59db953/nvidia_cuda_cupti_cu12-12.6.80-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:6768bad6cab4f19e8292125e5f1ac8aa7d1718704012a0e3272a6f61c4bce132", size = 8917980, upload-time = "2024-11-20T17:36:04.019Z" },
{ url = "https://files.pythonhosted.org/packages/a5/24/120ee57b218d9952c379d1e026c4479c9ece9997a4fb46303611ee48f038/nvidia_cuda_cupti_cu12-12.6.80-py3-none-manylinux2014_x86_64.whl", hash = "sha256:a3eff6cdfcc6a4c35db968a06fcadb061cbc7d6dde548609a941ff8701b98b73", size = 8917972, upload-time = "2024-10-01T16:58:06.036Z" },
]
[[package]]
name = "nvidia-cuda-nvrtc-cu12"
version = "12.6.77"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/75/2e/46030320b5a80661e88039f59060d1790298b4718944a65a7f2aeda3d9e9/nvidia_cuda_nvrtc_cu12-12.6.77-py3-none-manylinux2014_x86_64.whl", hash = "sha256:35b0cc6ee3a9636d5409133e79273ce1f3fd087abb0532d2d2e8fff1fe9efc53", size = 23650380, upload-time = "2024-10-01T17:00:14.643Z" },
]
[[package]]
name = "nvidia-cuda-runtime-cu12"
version = "12.6.77"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/e1/23/e717c5ac26d26cf39a27fbc076240fad2e3b817e5889d671b67f4f9f49c5/nvidia_cuda_runtime_cu12-12.6.77-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:ba3b56a4f896141e25e19ab287cd71e52a6a0f4b29d0d31609f60e3b4d5219b7", size = 897690, upload-time = "2024-11-20T17:35:30.697Z" },
{ url = "https://files.pythonhosted.org/packages/f0/62/65c05e161eeddbafeca24dc461f47de550d9fa8a7e04eb213e32b55cfd99/nvidia_cuda_runtime_cu12-12.6.77-py3-none-manylinux2014_x86_64.whl", hash = "sha256:a84d15d5e1da416dd4774cb42edf5e954a3e60cc945698dc1d5be02321c44dc8", size = 897678, upload-time = "2024-10-01T16:57:33.821Z" },
]
[[package]]
name = "nvidia-cudnn-cu12"
version = "9.5.1.17"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [ dependencies = [
{ name = "nvidia-cublas-cu12", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, { name = "nvidia-cuda-nvrtc", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
] ]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/2a/78/4535c9c7f859a64781e43c969a3a7e84c54634e319a996d43ef32ce46f83/nvidia_cudnn_cu12-9.5.1.17-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:30ac3869f6db17d170e0e556dd6cc5eee02647abc31ca856634d5a40f82c15b2", size = 570988386, upload-time = "2024-10-25T19:54:26.39Z" }, { url = "https://files.pythonhosted.org/packages/a7/a1/0bd24ee8c8d03adac032fd2909426a00c88f8c57961b1277ded97f91119f/nvidia_cublas-13.1.1.3-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:b7a210458267ac818974c53038fbec2e969d5c99f305ab15c72522fa9f001dd5", size = 542848918, upload-time = "2026-04-08T18:46:22.985Z" },
{ url = "https://files.pythonhosted.org/packages/3b/cd/154ca20c38269e05eff77c1464e6c1da89f50a6390b565e9d82e06bc11e1/nvidia_cublas-13.1.1.3-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:37936a16db8fe4ac1f065c2139360608a543a09275cb1a1af612e08cfa065436", size = 423138758, upload-time = "2026-04-08T18:46:58.655Z" },
] ]
[[package]] [[package]]
name = "nvidia-cufft-cu12" name = "nvidia-cuda-cupti"
version = "11.3.0.4" version = "13.0.85"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/2a/2a/80353b103fc20ce05ef51e928daed4b6015db4aaa9162ed0997090fe2250/nvidia_cuda_cupti-13.0.85-py3-none-manylinux_2_25_aarch64.whl", hash = "sha256:796bd679890ee55fb14a94629b698b6db54bcfd833d391d5e94017dd9d7d3151", size = 10310827, upload-time = "2025-09-04T08:26:42.012Z" },
{ url = "https://files.pythonhosted.org/packages/33/6d/737d164b4837a9bbd202f5ae3078975f0525a55730fe871d8ed4e3b952b0/nvidia_cuda_cupti-13.0.85-py3-none-manylinux_2_25_x86_64.whl", hash = "sha256:4eb01c08e859bf924d222250d2e8f8b8ff6d3db4721288cf35d14252a4d933c8", size = 10715597, upload-time = "2025-09-04T08:26:51.312Z" },
]
[[package]]
name = "nvidia-cuda-nvrtc"
version = "13.0.88"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c3/68/483a78f5e8f31b08fb1bb671559968c0ca3a065ac7acabfc7cee55214fd6/nvidia_cuda_nvrtc-13.0.88-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl", hash = "sha256:ad9b6d2ead2435f11cbb6868809d2adeeee302e9bb94bcf0539c7a40d80e8575", size = 90215200, upload-time = "2025-09-04T08:28:44.204Z" },
{ url = "https://files.pythonhosted.org/packages/b7/dc/6bb80850e0b7edd6588d560758f17e0550893a1feaf436807d64d2da040f/nvidia_cuda_nvrtc-13.0.88-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:d27f20a0ca67a4bb34268a5e951033496c5b74870b868bacd046b1b8e0c3267b", size = 43015449, upload-time = "2025-09-04T08:28:20.239Z" },
]
[[package]]
name = "nvidia-cuda-runtime"
version = "13.0.96"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/87/4f/17d7b9b8e285199c58ce28e31b5c5bbaa4d8271af06a89b6405258245de2/nvidia_cuda_runtime-13.0.96-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:ef9bcbe90493a2b9d810e43d249adb3d02e98dd30200d86607d8d02687c43f55", size = 2261060, upload-time = "2025-10-09T08:55:15.78Z" },
{ url = "https://files.pythonhosted.org/packages/2e/24/d1558f3b68b1d26e706813b1d10aa1d785e4698c425af8db8edc3dced472/nvidia_cuda_runtime-13.0.96-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:7f82250d7782aa23b6cfe765ecc7db554bd3c2870c43f3d1821f1d18aebf0548", size = 2243632, upload-time = "2025-10-09T08:55:36.117Z" },
]
[[package]]
name = "nvidia-cudnn-cu13"
version = "9.20.0.48"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [ dependencies = [
{ name = "nvidia-nvjitlink-cu12", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, { name = "nvidia-cublas", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
] ]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/8f/16/73727675941ab8e6ffd86ca3a4b7b47065edcca7a997920b831f8147c99d/nvidia_cufft_cu12-11.3.0.4-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:ccba62eb9cef5559abd5e0d54ceed2d9934030f51163df018532142a8ec533e5", size = 200221632, upload-time = "2024-11-20T17:41:32.357Z" }, { url = "https://files.pythonhosted.org/packages/56/c5/83384d846b2fd17c44bd499b36c75a45ed4f095fbbb2252294e89cea5c5c/nvidia_cudnn_cu13-9.20.0.48-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:e31454ae00094b0c55319d9d15b6fa2fc50a9e1c0f5c8c80fb75258234e731e1", size = 444574296, upload-time = "2026-03-09T19:28:27.751Z" },
{ url = "https://files.pythonhosted.org/packages/60/de/99ec247a07ea40c969d904fc14f3a356b3e2a704121675b75c366b694ee1/nvidia_cufft_cu12-11.3.0.4-py3-none-manylinux2014_x86_64.whl", hash = "sha256:768160ac89f6f7b459bee747e8d175dbf53619cfe74b2a5636264163138013ca", size = 200221622, upload-time = "2024-10-01T17:03:58.79Z" }, { url = "https://files.pythonhosted.org/packages/6e/5e/edb9c0ae051602c3ccaffe424256463636d639e27d7f302dde9975ef9e7a/nvidia_cudnn_cu13-9.20.0.48-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:0c45dd8eeb50b603f07995b1b300c62ffe6a1980482b82b3bcf94a4ca9d49304", size = 366173588, upload-time = "2026-03-09T19:29:34.474Z" },
] ]
[[package]] [[package]]
name = "nvidia-cufile-cu12" name = "nvidia-cufft"
version = "1.11.1.6" version = "12.0.0.61"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/b2/66/cc9876340ac68ae71b15c743ddb13f8b30d5244af344ec8322b449e35426/nvidia_cufile_cu12-1.11.1.6-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:cc23469d1c7e52ce6c1d55253273d32c565dd22068647f3aa59b3c6b005bf159", size = 1142103, upload-time = "2024-11-20T17:42:11.83Z" },
]
[[package]]
name = "nvidia-curand-cu12"
version = "10.3.7.77"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/73/1b/44a01c4e70933637c93e6e1a8063d1e998b50213a6b65ac5a9169c47e98e/nvidia_curand_cu12-10.3.7.77-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:a42cd1344297f70b9e39a1e4f467a4e1c10f1da54ff7a85c12197f6c652c8bdf", size = 56279010, upload-time = "2024-11-20T17:42:50.958Z" },
{ url = "https://files.pythonhosted.org/packages/4a/aa/2c7ff0b5ee02eaef890c0ce7d4f74bc30901871c5e45dee1ae6d0083cd80/nvidia_curand_cu12-10.3.7.77-py3-none-manylinux2014_x86_64.whl", hash = "sha256:99f1a32f1ac2bd134897fc7a203f779303261268a65762a623bf30cc9fe79117", size = 56279000, upload-time = "2024-10-01T17:04:45.274Z" },
]
[[package]]
name = "nvidia-cusolver-cu12"
version = "11.7.1.2"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [ dependencies = [
{ name = "nvidia-cublas-cu12", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
{ name = "nvidia-cusparse-cu12", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
{ name = "nvidia-nvjitlink-cu12", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
] ]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/f0/6e/c2cf12c9ff8b872e92b4a5740701e51ff17689c4d726fca91875b07f655d/nvidia_cusolver_cu12-11.7.1.2-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:e9e49843a7707e42022babb9bcfa33c29857a93b88020c4e4434656a655b698c", size = 158229790, upload-time = "2024-11-20T17:43:43.211Z" }, { url = "https://files.pythonhosted.org/packages/8b/ae/f417a75c0259e85c1d2f83ca4e960289a5f814ed0cea74d18c353d3e989d/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:2708c852ef8cd89d1d2068bdbece0aa188813a0c934db3779b9b1faa8442e5f5", size = 214053554, upload-time = "2025-09-04T08:31:38.196Z" },
{ url = "https://files.pythonhosted.org/packages/9f/81/baba53585da791d043c10084cf9553e074548408e04ae884cfe9193bd484/nvidia_cusolver_cu12-11.7.1.2-py3-none-manylinux2014_x86_64.whl", hash = "sha256:6cf28f17f64107a0c4d7802be5ff5537b2130bfc112f25d5a30df227058ca0e6", size = 158229780, upload-time = "2024-10-01T17:05:39.875Z" }, { url = "https://files.pythonhosted.org/packages/a8/2f/7b57e29836ea8714f81e9898409196f47d772d5ddedddf1592eadb8ab743/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:6c44f692dce8fd5ffd3e3df134b6cdb9c2f72d99cf40b62c32dde45eea9ddad3", size = 214085489, upload-time = "2025-09-04T08:31:56.044Z" },
] ]
[[package]] [[package]]
name = "nvidia-cusparse-cu12" name = "nvidia-cufile"
version = "12.5.4.2" version = "1.15.1.6"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/3f/70/4f193de89a48b71714e74602ee14d04e4019ad36a5a9f20c425776e72cd6/nvidia_cufile-1.15.1.6-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:08a3ecefae5a01c7f5117351c64f17c7c62efa5fffdbe24fc7d298da19cd0b44", size = 1223672, upload-time = "2025-09-04T08:32:22.779Z" },
{ url = "https://files.pythonhosted.org/packages/ab/73/cc4a14c9813a8a0d509417cf5f4bdaba76e924d58beb9864f5a7baceefbf/nvidia_cufile-1.15.1.6-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:bdc0deedc61f548bddf7733bdc216456c2fdb101d020e1ab4b88d232d5e2f6d1", size = 1136992, upload-time = "2025-09-04T08:32:14.119Z" },
]
[[package]]
name = "nvidia-curand"
version = "10.4.0.35"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1e/72/7c2ae24fb6b63a32e6ae5d241cc65263ea18d08802aaae087d9f013335a2/nvidia_curand-10.4.0.35-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:133df5a7509c3e292aaa2b477afd0194f06ce4ea24d714d616ff36439cee349a", size = 61962106, upload-time = "2025-08-04T10:21:41.128Z" },
{ url = "https://files.pythonhosted.org/packages/a5/9f/be0a41ca4a4917abf5cb9ae0daff1a6060cc5de950aec0396de9f3b52bc5/nvidia_curand-10.4.0.35-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:1aee33a5da6e1db083fe2b90082def8915f30f3248d5896bcec36a579d941bfc", size = 59544258, upload-time = "2025-08-04T10:22:03.992Z" },
]
[[package]]
name = "nvidia-cusolver"
version = "12.0.4.66"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [ dependencies = [
{ name = "nvidia-nvjitlink-cu12", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, { name = "nvidia-cublas", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
{ name = "nvidia-cusparse", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
{ name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
] ]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/06/1e/b8b7c2f4099a37b96af5c9bb158632ea9e5d9d27d7391d7eb8fc45236674/nvidia_cusparse_cu12-12.5.4.2-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:7556d9eca156e18184b94947ade0fba5bb47d69cec46bf8660fd2c71a4b48b73", size = 216561367, upload-time = "2024-11-20T17:44:54.824Z" }, { url = "https://files.pythonhosted.org/packages/c8/c3/b30c9e935fc01e3da443ec0116ed1b2a009bb867f5324d3f2d7e533e776b/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:02c2457eaa9e39de20f880f4bd8820e6a1cfb9f9a34f820eb12a155aa5bc92d2", size = 223467760, upload-time = "2025-09-04T08:33:04.222Z" },
{ url = "https://files.pythonhosted.org/packages/43/ac/64c4316ba163e8217a99680c7605f779accffc6a4bcd0c778c12948d3707/nvidia_cusparse_cu12-12.5.4.2-py3-none-manylinux2014_x86_64.whl", hash = "sha256:23749a6571191a215cb74d1cdbff4a86e7b19f1200c071b3fcf844a5bea23a2f", size = 216561357, upload-time = "2024-10-01T17:06:29.861Z" }, { url = "https://files.pythonhosted.org/packages/5f/67/cba3777620cdacb99102da4042883709c41c709f4b6323c10781a9c3aa34/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:0a759da5dea5c0ea10fd307de75cdeb59e7ea4fcb8add0924859b944babf1112", size = 200941980, upload-time = "2025-09-04T08:33:22.767Z" },
] ]
[[package]] [[package]]
name = "nvidia-cusparselt-cu12" name = "nvidia-cusparse"
version = "0.6.3" version = "12.6.3.3"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/3b/9a/72ef35b399b0e183bc2e8f6f558036922d453c4d8237dab26c666a04244b/nvidia_cusparselt_cu12-0.6.3-py3-none-manylinux2014_x86_64.whl", hash = "sha256:e5c8a26c36445dd2e6812f1177978a24e2d37cacce7e090f297a688d1ec44f46", size = 156785796, upload-time = "2024-10-15T21:29:17.709Z" }, { url = "https://files.pythonhosted.org/packages/f8/94/5c26f33738ae35276672f12615a64bd008ed5be6d1ebcb23579285d960a9/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:80bcc4662f23f1054ee334a15c72b8940402975e0eab63178fc7e670aa59472c", size = 162155568, upload-time = "2025-09-04T08:33:42.864Z" },
{ url = "https://files.pythonhosted.org/packages/fa/18/623c77619c31d62efd55302939756966f3ecc8d724a14dab2b75f1508850/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:2b3c89c88d01ee0e477cb7f82ef60a11a4bcd57b6b87c33f789350b59759360b", size = 145942937, upload-time = "2025-09-04T08:33:58.029Z" },
] ]
[[package]] [[package]]
name = "nvidia-nccl-cu12" name = "nvidia-cusparselt-cu13"
version = "2.26.2" version = "0.8.1"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/67/ca/f42388aed0fddd64ade7493dbba36e1f534d4e6fdbdd355c6a90030ae028/nvidia_nccl_cu12-2.26.2-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:694cf3879a206553cc9d7dbda76b13efaf610fdb70a50cba303de1b0d1530ac6", size = 201319755, upload-time = "2025-03-13T00:29:55.296Z" }, { url = "https://files.pythonhosted.org/packages/46/e1/cdc1797eadf82d3a9a575a19b33fdc871a97edbec42c00b5b5e914f4aff4/nvidia_cusparselt_cu13-0.8.1-py3-none-manylinux2014_aarch64.whl", hash = "sha256:4dca476c50bf4780d46cd0bfbd82e2bc10a08e4fef7950917ce8d7578d22a23f", size = 221051344, upload-time = "2025-09-05T18:49:51.289Z" },
{ url = "https://files.pythonhosted.org/packages/34/7d/2661f2fb3ac4302f3a246f5fc030213ac60c1fe0bce84f9783dbd831dbb7/nvidia_cusparselt_cu13-0.8.1-py3-none-manylinux2014_x86_64.whl", hash = "sha256:786ce87568c303fadb5afcc7102d454cd3040d75f6f8626f5db460d1871f4dd0", size = 170148586, upload-time = "2025-09-05T18:50:50.248Z" },
] ]
[[package]] [[package]]
name = "nvidia-nvjitlink-cu12" name = "nvidia-nccl-cu13"
version = "12.6.85" version = "2.29.7"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/9d/d7/c5383e47c7e9bf1c99d5bd2a8c935af2b6d705ad831a7ec5c97db4d82f4f/nvidia_nvjitlink_cu12-12.6.85-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl", hash = "sha256:eedc36df9e88b682efe4309aa16b5b4e78c2407eac59e8c10a6a47535164369a", size = 19744971, upload-time = "2024-11-20T17:46:53.366Z" }, { url = "https://files.pythonhosted.org/packages/72/0d/daf50d44177ee0cbc7ff0a0c91eb5ff676c82be42f9a970bc7597f440c3a/nvidia_nccl_cu13-2.29.7-py3-none-manylinux_2_18_aarch64.whl", hash = "sha256:674a12383e3c38a1bcccae7d4f3633b37852230b6047883cb2f4c2d1b36d9bf5", size = 206014712, upload-time = "2026-03-03T05:34:20.843Z" },
{ url = "https://files.pythonhosted.org/packages/67/f4/58e4e91b6919367c7aafb8e36fce9aad1a3047e536bf7e2fd560927d3a4c/nvidia_nccl_cu13-2.29.7-py3-none-manylinux_2_18_x86_64.whl", hash = "sha256:edd81538446786ec3b73972543e53bb43bcaf0bfc8ef76cb679fcc390ffe136d", size = 205976000, upload-time = "2026-03-03T05:36:24.472Z" },
] ]
[[package]] [[package]]
name = "nvidia-nvtx-cu12" name = "nvidia-nvjitlink"
version = "12.6.77" version = "13.3.33"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/56/9a/fff8376f8e3d084cd1530e1ef7b879bb7d6d265620c95c1b322725c694f4/nvidia_nvtx_cu12-12.6.77-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:b90bed3df379fa79afbd21be8e04a0314336b8ae16768b58f2d34cb1d04cd7d2", size = 89276, upload-time = "2024-11-20T17:38:27.621Z" }, { url = "https://files.pythonhosted.org/packages/f0/ee/580ca6f29dcab0221db8706badca1bbbb084f1975c4d4e83329c3a7e31f0/nvidia_nvjitlink-13.3.33-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl", hash = "sha256:26a6de7fb4c8fdaa7703d3dad720d6d427ddfea5c48a528fd97c11733ad830e5", size = 40742423, upload-time = "2026-05-26T16:54:51.613Z" },
{ url = "https://files.pythonhosted.org/packages/9e/4e/0d0c945463719429b7bd21dece907ad0bde437a2ff12b9b12fee94722ab0/nvidia_nvtx_cu12-12.6.77-py3-none-manylinux2014_x86_64.whl", hash = "sha256:6574241a3ec5fdc9334353ab8c479fe75841dbe8f4532a8fc97ce63503330ba1", size = 89265, upload-time = "2024-10-01T17:00:38.172Z" }, { url = "https://files.pythonhosted.org/packages/69/30/45414e35ff2eee7db3da037e5707037ccf9d2b5218ffbdb055ea4d5aa98a/nvidia_nvjitlink-13.3.33-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:ce48b37dfeb3cb1eae4cf85adacb47d7a6539ea2272870c9a3628ce275c2037e", size = 39168635, upload-time = "2026-05-26T16:54:13.906Z" },
]
[[package]]
name = "nvidia-nvshmem-cu13"
version = "3.4.5"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/dc/0f/05cc9c720236dcd2db9c1ab97fff629e96821be2e63103569da0c9b72f19/nvidia_nvshmem_cu13-3.4.5-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:6dc2a197f38e5d0376ad52cd1a2a3617d3cdc150fd5966f4aee9bcebb1d68fe9", size = 60215947, upload-time = "2025-09-06T00:32:20.022Z" },
{ url = "https://files.pythonhosted.org/packages/3c/35/a9bf80a609e74e3b000fef598933235c908fcefcef9026042b8e6dfde2a9/nvidia_nvshmem_cu13-3.4.5-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:290f0a2ee94c9f3687a02502f3b9299a9f9fe826e6d0287ee18482e78d495b80", size = 60412546, upload-time = "2025-09-06T00:32:41.564Z" },
]
[[package]]
name = "nvidia-nvtx"
version = "13.0.85"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c2/f3/d86c845465a2723ad7e1e5c36dcd75ddb82898b3f53be47ebd429fb2fa5d/nvidia_nvtx-13.0.85-py3-none-manylinux1_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:4936d1d6780fbe68db454f5e72a42ff64d1fd6397df9f363ae786930fd5c1cd4", size = 148047, upload-time = "2025-09-04T08:29:01.761Z" },
{ url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878, upload-time = "2025-09-04T08:28:53.627Z" },
] ]
[[package]] [[package]]
@ -3865,49 +3961,45 @@ wheels = [
[[package]] [[package]]
name = "torch" name = "torch"
version = "2.7.1" version = "2.13.0"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [ dependencies = [
{ name = "cuda-bindings", marker = "python_full_version < '3.15' and sys_platform == 'linux'" },
{ name = "cuda-toolkit", extra = ["cublas", "cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"], marker = "sys_platform == 'linux'" },
{ name = "filelock" }, { name = "filelock" },
{ name = "fsspec" }, { name = "fsspec" },
{ name = "jinja2" }, { name = "jinja2" },
{ name = "networkx" }, { name = "networkx" },
{ name = "nvidia-cublas-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" }, { name = "nvidia-cudnn-cu13", marker = "sys_platform == 'linux'" },
{ name = "nvidia-cuda-cupti-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" }, { name = "nvidia-cusparselt-cu13", marker = "sys_platform == 'linux'" },
{ name = "nvidia-cuda-nvrtc-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" }, { name = "nvidia-nccl-cu13", marker = "sys_platform == 'linux'" },
{ name = "nvidia-cuda-runtime-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" }, { name = "nvidia-nvshmem-cu13", marker = "sys_platform == 'linux'" },
{ name = "nvidia-cudnn-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" }, { name = "setuptools" },
{ name = "nvidia-cufft-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cufile-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-curand-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cusolver-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cusparse-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cusparselt-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-nccl-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-nvjitlink-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-nvtx-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "setuptools", marker = "python_full_version >= '3.12'" },
{ name = "sympy" }, { name = "sympy" },
{ name = "triton", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" }, { name = "triton", marker = "python_full_version < '3.15' and sys_platform == 'linux'" },
{ name = "typing-extensions" }, { name = "typing-extensions" },
] ]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/11/56/2eae3494e3d375533034a8e8cf0ba163363e996d85f0629441fa9d9843fe/torch-2.7.1-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:236f501f2e383f1cb861337bdf057712182f910f10aeaf509065d54d339e49b2", size = 99093039, upload-time = "2025-06-04T17:39:06.963Z" }, { url = "https://files.pythonhosted.org/packages/5b/fe/cba54dc58523434919b66f13a667e36e436deddd77ca519e96553617d4ec/torch-2.13.0-cp311-cp311-macosx_14_0_arm64.whl", hash = "sha256:e76f9bcecc52b8ff711239a2f7547d5353df95878ab232f0773c1d95928b92f8", size = 111187938, upload-time = "2026-07-08T16:05:17.065Z" },
{ url = "https://files.pythonhosted.org/packages/e5/94/34b80bd172d0072c9979708ccd279c2da2f55c3ef318eceec276ab9544a4/torch-2.7.1-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:06eea61f859436622e78dd0cdd51dbc8f8c6d76917a9cf0555a333f9eac31ec1", size = 821174704, upload-time = "2025-06-04T17:37:03.799Z" }, { url = "https://files.pythonhosted.org/packages/c2/59/1e3160e18e12aa3038390efab3ce02b36a9d4d6a527ecdd8520dca2e68d8/torch-2.13.0-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:092790c696a760c729fd5722835f50b9d81fd7c8f141571f3f3cf4081a8f664c", size = 427199369, upload-time = "2026-07-08T16:04:51.054Z" },
{ url = "https://files.pythonhosted.org/packages/50/9e/acf04ff375b0b49a45511c55d188bcea5c942da2aaf293096676110086d1/torch-2.7.1-cp311-cp311-win_amd64.whl", hash = "sha256:8273145a2e0a3c6f9fd2ac36762d6ee89c26d430e612b95a99885df083b04e52", size = 216095937, upload-time = "2025-06-04T17:39:24.83Z" }, { url = "https://files.pythonhosted.org/packages/01/79/1f2d34ad7034ee1c7ffc1cf8bf0f8213af2a81df6ecdb3997ecec107c09d/torch-2.13.0-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:60fcdcb2f3876e21146cb4524ef06397d727ca9ad5f020818547e25075fe3cb7", size = 526574961, upload-time = "2026-07-08T16:04:07.075Z" },
{ url = "https://files.pythonhosted.org/packages/5b/2b/d36d57c66ff031f93b4fa432e86802f84991477e522adcdffd314454326b/torch-2.7.1-cp311-none-macosx_11_0_arm64.whl", hash = "sha256:aea4fc1bf433d12843eb2c6b2204861f43d8364597697074c8d38ae2507f8730", size = 68640034, upload-time = "2025-06-04T17:39:17.989Z" }, { url = "https://files.pythonhosted.org/packages/6c/fd/0f2ce40f58aefbdb3392f9acce3c8171940943ae2d661f70558bfa73befb/torch-2.13.0-cp311-cp311-win_amd64.whl", hash = "sha256:a0d8b11f16a48d60e2015d8213aa0390744cbebb98e58b62b3514dddc656e330", size = 122015870, upload-time = "2026-07-08T16:05:27.59Z" },
{ url = "https://files.pythonhosted.org/packages/87/93/fb505a5022a2e908d81fe9a5e0aa84c86c0d5f408173be71c6018836f34e/torch-2.7.1-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:27ea1e518df4c9de73af7e8a720770f3628e7f667280bce2be7a16292697e3fa", size = 98948276, upload-time = "2025-06-04T17:39:12.852Z" }, { url = "https://files.pythonhosted.org/packages/c4/3a/ed0f4d4d1dcde03bced7aac9a28e800abcdc0cbd06b6775044c9fbd877b7/torch-2.13.0-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:2fe228aba290d14b9f31b049be550dbd469c3fd3013d7a19705b30454da97027", size = 111213045, upload-time = "2026-07-08T16:05:22.997Z" },
{ url = "https://files.pythonhosted.org/packages/56/7e/67c3fe2b8c33f40af06326a3d6ae7776b3e3a01daa8f71d125d78594d874/torch-2.7.1-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:c33360cfc2edd976c2633b3b66c769bdcbbf0e0b6550606d188431c81e7dd1fc", size = 821025792, upload-time = "2025-06-04T17:34:58.747Z" }, { url = "https://files.pythonhosted.org/packages/df/a9/f6a2a4d763ff1df02e9a64c477029db614295bc9367f4131223791ccc243/torch-2.13.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:572df8be8ffb4599c88cbd6a0726f1f854f4da65d2e3c09f0e2c2283333cd6d4", size = 427210998, upload-time = "2026-07-08T16:04:37.708Z" },
{ url = "https://files.pythonhosted.org/packages/a1/37/a37495502bc7a23bf34f89584fa5a78e25bae7b8da513bc1b8f97afb7009/torch-2.7.1-cp312-cp312-win_amd64.whl", hash = "sha256:d8bf6e1856ddd1807e79dc57e54d3335f2b62e6f316ed13ed3ecfe1fc1df3d8b", size = 216050349, upload-time = "2025-06-04T17:38:59.709Z" }, { url = "https://files.pythonhosted.org/packages/f3/82/fea946351658e6534db52d2cc12bc53087cbf87f9440c5f180f367c1950b/torch-2.13.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:796633c4cdf0fe2cdced72d8f88f22e73dbcfce83132763162f6d4bff13b820b", size = 526605292, upload-time = "2026-07-08T16:04:22.81Z" },
{ url = "https://files.pythonhosted.org/packages/3a/60/04b77281c730bb13460628e518c52721257814ac6c298acd25757f6a175c/torch-2.7.1-cp312-none-macosx_11_0_arm64.whl", hash = "sha256:787687087412c4bd68d315e39bc1223f08aae1d16a9e9771d95eabbb04ae98fb", size = 68645146, upload-time = "2025-06-04T17:38:52.97Z" }, { url = "https://files.pythonhosted.org/packages/21/d6/e8f3c6f7e01f626f77259de9860d2a78bc84c40539e28e79b7e98b0bb659/torch-2.13.0-cp312-cp312-win_amd64.whl", hash = "sha256:024c6cc0c1b085f2f91f20a3dc27b0471d021c31ce84b81be3afdc39f791fd9d", size = 122057313, upload-time = "2026-07-08T16:03:53.43Z" },
{ url = "https://files.pythonhosted.org/packages/66/81/e48c9edb655ee8eb8c2a6026abdb6f8d2146abd1f150979ede807bb75dcb/torch-2.7.1-cp313-cp313-manylinux_2_28_aarch64.whl", hash = "sha256:03563603d931e70722dce0e11999d53aa80a375a3d78e6b39b9f6805ea0a8d28", size = 98946649, upload-time = "2025-06-04T17:38:43.031Z" }, { url = "https://files.pythonhosted.org/packages/0d/fa/c1c10b7aff4a9a3e8956d4f0a5f468fa6db7abc3208805719076772b4833/torch-2.13.0-cp313-cp313-macosx_14_0_arm64.whl", hash = "sha256:33449899ce5496c1b84b4853179d94fd102028ae1407314d9fb956bb79e70d09", size = 111213743, upload-time = "2026-07-08T16:03:28.579Z" },
{ url = "https://files.pythonhosted.org/packages/3a/24/efe2f520d75274fc06b695c616415a1e8a1021d87a13c68ff9dce733d088/torch-2.7.1-cp313-cp313-manylinux_2_28_x86_64.whl", hash = "sha256:d632f5417b6980f61404a125b999ca6ebd0b8b4bbdbb5fbbba44374ab619a412", size = 821033192, upload-time = "2025-06-04T17:38:09.146Z" }, { url = "https://files.pythonhosted.org/packages/11/18/9ecb37b56293a0be8d80f810bf672a72fe7e02f8b475d5ef1b9bf8a0d748/torch-2.13.0-cp313-cp313-manylinux_2_28_aarch64.whl", hash = "sha256:1e09d6a722504957c694faceca843acde562786df1144ebcc5a74075ec7f6005", size = 427213008, upload-time = "2026-07-08T16:03:44.106Z" },
{ url = "https://files.pythonhosted.org/packages/dd/d9/9c24d230333ff4e9b6807274f6f8d52a864210b52ec794c5def7925f4495/torch-2.7.1-cp313-cp313-win_amd64.whl", hash = "sha256:23660443e13995ee93e3d844786701ea4ca69f337027b05182f5ba053ce43b38", size = 216055668, upload-time = "2025-06-04T17:38:36.253Z" }, { url = "https://files.pythonhosted.org/packages/d4/5a/7c50ba1b7b713d71d34669c6d13dab0a11531a3eceb0307a5162dbfec0f7/torch-2.13.0-cp313-cp313-manylinux_2_28_x86_64.whl", hash = "sha256:a3a9a21312872af8a26950b2c15680335a386a1f56ed03e780653d78b9607e9e", size = 526602329, upload-time = "2026-07-08T16:03:12.649Z" },
{ url = "https://files.pythonhosted.org/packages/95/bf/e086ee36ddcef9299f6e708d3b6c8487c1651787bb9ee2939eb2a7f74911/torch-2.7.1-cp313-cp313t-macosx_14_0_arm64.whl", hash = "sha256:0da4f4dba9f65d0d203794e619fe7ca3247a55ffdcbd17ae8fb83c8b2dc9b585", size = 68925988, upload-time = "2025-06-04T17:38:29.273Z" }, { url = "https://files.pythonhosted.org/packages/91/3d/e7adcc6aaf36961cd18f56cf8ad0f3058c3a5c84ccf391762176c94581b8/torch-2.13.0-cp313-cp313-win_amd64.whl", hash = "sha256:49b58f1e2c52440abb6f17c28f0335fe6c6d01ad1a7f55b0183b81e4b34d64e6", size = 122057920, upload-time = "2026-07-08T16:03:01.808Z" },
{ url = "https://files.pythonhosted.org/packages/69/6a/67090dcfe1cf9048448b31555af6efb149f7afa0a310a366adbdada32105/torch-2.7.1-cp313-cp313t-manylinux_2_28_aarch64.whl", hash = "sha256:e08d7e6f21a617fe38eeb46dd2213ded43f27c072e9165dc27300c9ef9570934", size = 99028857, upload-time = "2025-06-04T17:37:50.956Z" }, { url = "https://files.pythonhosted.org/packages/36/76/6dcc7f0c07052102dd36f83cbc5800842a909c8c3fbf1a7f8a5844954de9/torch-2.13.0-cp314-cp314-macosx_14_0_arm64.whl", hash = "sha256:d849b390e07d8d333ce8ecaf91b273c656c598379a19c9acf1318a883f6b391c", size = 111227066, upload-time = "2026-07-08T16:03:33.6Z" },
{ url = "https://files.pythonhosted.org/packages/90/1c/48b988870823d1cc381f15ec4e70ed3d65e043f43f919329b0045ae83529/torch-2.7.1-cp313-cp313t-manylinux_2_28_x86_64.whl", hash = "sha256:30207f672328a42df4f2174b8f426f354b2baa0b7cca3a0adb3d6ab5daf00dc8", size = 821098066, upload-time = "2025-06-04T17:37:33.939Z" }, { url = "https://files.pythonhosted.org/packages/e9/09/2c10e8cd0e00fa5d23c052df6ce467eaa7182399f5e0f824f1e4ff42ccae/torch-2.13.0-cp314-cp314-manylinux_2_28_aarch64.whl", hash = "sha256:a3893dc2da0a972a8ca5d698c85a9f967559ac5f8ee1797b77408aa8734d073c", size = 427226309, upload-time = "2026-07-08T16:02:53.127Z" },
{ url = "https://files.pythonhosted.org/packages/7b/eb/10050d61c9d5140c5dc04a89ed3257ef1a6b93e49dd91b95363d757071e0/torch-2.7.1-cp313-cp313t-win_amd64.whl", hash = "sha256:79042feca1c634aaf6603fe6feea8c6b30dfa140a6bbc0b973e2260c7e79a22e", size = 216336310, upload-time = "2025-06-04T17:36:09.862Z" }, { url = "https://files.pythonhosted.org/packages/76/c6/22c2102bbef14ca6a6cb4c20e42f088e49c5f812be4e160ae57502e325f9/torch-2.13.0-cp314-cp314-manylinux_2_28_x86_64.whl", hash = "sha256:49f1ea385c754e54919408a9bb3b5a72b0b755bbe2c916c1d6f70afbec4908a2", size = 526614507, upload-time = "2026-07-08T16:02:16.441Z" },
{ url = "https://files.pythonhosted.org/packages/b1/29/beb45cdf5c4fc3ebe282bf5eafc8dfd925ead7299b3c97491900fe5ed844/torch-2.7.1-cp313-none-macosx_11_0_arm64.whl", hash = "sha256:988b0cbc4333618a1056d2ebad9eb10089637b659eb645434d0809d8d937b946", size = 68645708, upload-time = "2025-06-04T17:34:39.852Z" }, { url = "https://files.pythonhosted.org/packages/2b/0c/7d1deb6bce5bc3e6042caf39100ac768eba3b9a098e1dddd16f75bd6489b/torch-2.13.0-cp314-cp314-win_amd64.whl", hash = "sha256:4f8573e3ce9ebcd53fe922f01077a6085ccdfbe5f12fd215883a9d87d7a744fd", size = 122051871, upload-time = "2026-07-08T16:03:23.521Z" },
{ url = "https://files.pythonhosted.org/packages/f4/ce/aa8b7f9949d32e0f2f624f342bc3b48112c1b8a130288465938bc83bcbf9/torch-2.13.0-cp314-cp314t-macosx_14_0_arm64.whl", hash = "sha256:c28def70706c2f9ecc752574766e8ae4da9b810ab6676b611166761a78a9f1e1", size = 111537025, upload-time = "2026-07-08T16:02:44.28Z" },
{ url = "https://files.pythonhosted.org/packages/69/d1/491e3a0389430946145888b0203f2b6a759ce2a61481b96a85c2da4f2ced/torch-2.13.0-cp314-cp314t-manylinux_2_28_aarch64.whl", hash = "sha256:31061ff56ed8fbf26c749806905aeb749ebeb819810fd5d52508aa5afd90dddc", size = 427219769, upload-time = "2026-07-08T16:02:31.18Z" },
{ url = "https://files.pythonhosted.org/packages/9a/1d/38006e045bf0a1fc28ef01e757c554e59e59a8770c284bc4f47b14e60441/torch-2.13.0-cp314-cp314t-manylinux_2_28_x86_64.whl", hash = "sha256:cc26eead4cf51d0b544e31e364dcf000846549c273bd148936fe9d24d29acb92", size = 526571320, upload-time = "2026-07-08T16:01:59.348Z" },
{ url = "https://files.pythonhosted.org/packages/56/94/655c91992a882bd5071aa0b5d22a07dbb130d801e872be97c0b627a7c693/torch-2.13.0-cp314-cp314t-win_amd64.whl", hash = "sha256:a7de8a313090dc5c7d7ba4bfe5c3be222528f9a4dba1acc83bddb1157360c4b8", size = 122306773, upload-time = "2026-07-08T16:02:39.832Z" },
] ]
[[package]] [[package]]
@ -3950,16 +4042,19 @@ wheels = [
[[package]] [[package]]
name = "triton" name = "triton"
version = "3.3.1" version = "3.7.1"
source = { registry = "https://pypi.org/simple" } source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "setuptools", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
]
wheels = [ wheels = [
{ url = "https://files.pythonhosted.org/packages/21/2f/3e56ea7b58f80ff68899b1dbe810ff257c9d177d288c6b0f55bf2fe4eb50/triton-3.3.1-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b31e3aa26f8cb3cc5bf4e187bf737cbacf17311e1112b781d4a059353dfd731b", size = 155689937, upload-time = "2025-05-29T23:39:44.182Z" }, { url = "https://files.pythonhosted.org/packages/7b/f9/19d842d06a08559534fa1eaab6ca551b1bcf40f06620bddec1babaa2772d/triton-3.7.1-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d4a0e1cd4c4a76370ed74a8432a53cea28716827d19e40ffc732233e35ceb3f6", size = 184664887, upload-time = "2026-06-17T20:03:42.913Z" },
{ url = "https://files.pythonhosted.org/packages/24/5f/950fb373bf9c01ad4eb5a8cd5eaf32cdf9e238c02f9293557a2129b9c4ac/triton-3.3.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9999e83aba21e1a78c1f36f21bce621b77bcaa530277a50484a7cb4a822f6e43", size = 155669138, upload-time = "2025-05-29T23:39:51.771Z" }, { url = "https://files.pythonhosted.org/packages/cd/5e/fce69606f7f240297f163e25539906732b199530d486ce67ae319877e821/triton-3.7.1-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:6744957e9fd610a29680ec2346057d0c86948ed3812468670719f391e94b44a5", size = 197701306, upload-time = "2026-06-17T19:53:13.673Z" },
{ url = "https://files.pythonhosted.org/packages/74/1f/dfb531f90a2d367d914adfee771babbd3f1a5b26c3f5fbc458dee21daa78/triton-3.3.1-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b89d846b5a4198317fec27a5d3a609ea96b6d557ff44b56c23176546023c4240", size = 155673035, upload-time = "2025-05-29T23:40:02.468Z" }, { url = "https://files.pythonhosted.org/packages/94/fa/f856e24deb462d5f18bd4b5a746957862ab9b6ee5834bda60605ec348366/triton-3.7.1-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9497f2e696ee368862a181a90b2dcc03ca978cc4f602abd67c7d81022a6988e1", size = 184692359, upload-time = "2026-06-17T20:03:48.288Z" },
{ url = "https://files.pythonhosted.org/packages/28/71/bd20ffcb7a64c753dc2463489a61bf69d531f308e390ad06390268c4ea04/triton-3.3.1-cp313-cp313t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:a3198adb9d78b77818a5388bff89fa72ff36f9da0bc689db2f0a651a67ce6a42", size = 155735832, upload-time = "2025-05-29T23:40:10.522Z" }, { url = "https://files.pythonhosted.org/packages/c4/6f/fb96d15db6f36d6eae4cafb998c2e0353bf59d7c4ea1662d7497f269134a/triton-3.7.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:7e40869937a68206ec70d7f25bb7ec6433cb083f9135e1f36dbd318dc449a728", size = 197719725, upload-time = "2026-06-17T19:53:20.419Z" },
{ url = "https://files.pythonhosted.org/packages/00/42/c5089d4d9327fcd1e862c599cc2927f39418f84dd11a84cb2ccff9d4787a/triton-3.7.1-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:cdbfc09d9ec58bc5e68321525653220de7515c199e7a8097a97c85e62b52cd0a", size = 184694629, upload-time = "2026-06-17T20:03:53.444Z" },
{ url = "https://files.pythonhosted.org/packages/07/42/2c3ac59253ae8892b6f307875263dd23dc875cdf732d3aea40d6d41fb7cb/triton-3.7.1-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:58c0e131da05134a2a4788ccbcc0c1105cf0f54c8e98f19e34cd465396dc15eb", size = 197729241, upload-time = "2026-06-17T19:53:27.801Z" },
{ url = "https://files.pythonhosted.org/packages/40/71/e01aa7ad573883ed9456f130226babdec70b005e098c4d6226a6238e761b/triton-3.7.1-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:fe4ea396a06171f1f1f58cbd39c70b09294398f7dd7c620939bab54ad6f934fa", size = 184705764, upload-time = "2026-06-17T20:03:59.064Z" },
{ url = "https://files.pythonhosted.org/packages/a4/09/5683146fda6a2b569deb78ccfd8fbfea8bfe55f726b081c0a6bb18dd6f28/triton-3.7.1-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2020153b08280415ec0da6607834e79166442147e78e144df06b508c75b186d2", size = 197729537, upload-time = "2026-06-17T19:53:35.516Z" },
{ url = "https://files.pythonhosted.org/packages/e9/f8/448220c3092019f9fdfab39ec47985968181d67da34b44f6a7f6280a5cbb/triton-3.7.1-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c58e4c61f0c73b5dba3b5d19b4a7093c32f90dc18b2a7f121a7c16ccd31107b7", size = 184814760, upload-time = "2026-06-17T20:04:04.984Z" },
{ url = "https://files.pythonhosted.org/packages/f0/ac/229b7d4589d2e5937310e72c6d46e89599d16a4a12b479ffa1499fee8eb8/triton-3.7.1-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:10ba85fa2cca4a2fbdeb36bf1cb082f2c252bda55bf9fccd74f65ec5bc647e68", size = 197824404, upload-time = "2026-06-17T19:53:42.772Z" },
] ]
[[package]] [[package]]