feat: more sweeping, data plane fixes

This commit is contained in:
Zachery Aaron Shores-Chmielewski 2026-07-27 12:48:37 +04:00
parent 292f1ea606
commit bc9ba691ac
28 changed files with 2541 additions and 114 deletions

View file

@ -1,6 +1,7 @@
/artifacts /artifacts
/data/raw /data/raw
/data/processed /data/processed
/data/airfrans_processed_full_50cases.tar.gz
/.venv /.venv
/notebooks /notebooks
__pycache__ __pycache__

View file

@ -4,25 +4,17 @@ seed = 20260723
artifact_dir = "artifacts/current_run/training_runs" artifact_dir = "artifacts/current_run/training_runs"
[data] [data]
root = "artifacts/data_cache/airfrans_streaming_processed/processed/full" root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900 train_cases = 900
val_cases = 50 val_cases = 50
test_cases = 50 test_cases = 50
points_per_case = 999999999 points_per_case = 999999999
batch_size = 4096 batch_size = 4096
source = "public_zip_streaming" source = "huggingface"
public_source_url = "https://data.isir.upmc.fr/extrality/NeurIPS_2022/OF_dataset.zip"
hf_repo_id = "zacheryasc/airfrans-processed" hf_repo_id = "zacheryasc/airfrans-processed"
hf_repo_type = "dataset" hf_repo_type = "dataset"
hf_path_prefix = "processed/full" hf_path_prefix = "processed/full"
cache_dir = "artifacts/data_cache/airfrans_streaming_processed/processed/full" cache_dir = "artifacts/data_cache/airfrans_processed_hf"
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] [model]
type = "film_fourier_inr" type = "film_fourier_inr"

View file

@ -0,0 +1,56 @@
[run]
name = "full_70gb_model_zoo_20260726_deeponet_branch_trunk"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/deeponet_branch_trunk"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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 = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "deeponet_branch_trunk"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -0,0 +1,56 @@
[run]
name = "full_70gb_model_zoo_20260726_film_fourier_inr"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/film_fourier_inr"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "film_fourier_inr"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -0,0 +1,56 @@
[run]
name = "full_70gb_model_zoo_20260726_meshgraphnet_or_point_transformer_local"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/meshgraphnet_or_point_transformer_local"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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 = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "meshgraphnet_or_point_transformer_local"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -0,0 +1,57 @@
[run]
name = "full_70gb_model_zoo_20260726_nerf_cfd_multires"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/nerf_cfd_multires"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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 = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "nerf_cfd_multires"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -0,0 +1,58 @@
[run]
name = "full_70gb_model_zoo_20260726_point_context_perceiver"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/point_context_perceiver"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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 = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "point_context_perceiver"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -0,0 +1,56 @@
[run]
name = "full_70gb_model_zoo_20260726_raster_fno_unet"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/raster_fno_unet"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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 = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "raster_fno_unet"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -0,0 +1,56 @@
[run]
name = "full_70gb_model_zoo_20260726_siren_conditioned_inr"
seed = 20260726
artifact_dir = "artifacts/current_run/training_runs/siren_conditioned_inr"
[data]
root = "artifacts/data_cache/airfrans_processed_hf/processed/full"
train_cases = 900
val_cases = 50
test_cases = 50
points_per_case = 999999999
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_hf"
[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 = 900
[stability]
max_grad_norm = 1.0
[observability]
backend = "wandb"
entity = "zacheryasc-personal"
project = "airfRANS-model-sweep"
group = "full_70gb_model_zoo_20260726"
tags = ["airfrans", "full-data-frontier", "70gb", "model-zoo", "hf-checkpoints", "siren_conditioned_inr"]
[huggingface]
enabled = true
repo_id = "zacheryasc/airfrans-frontier-checkpoints"
repo_type = "model"
path_prefix = "full_70gb_model_zoo_20260726"
private = false

View file

@ -3,6 +3,7 @@ name = "full_airfrans_incumbent_70gb_01"
timeout_minutes = 1440 timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs" local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5 max_attempts = 5
artifact_sync_interval_seconds = 60
[provider] [provider]
kind = "vastai" kind = "vastai"
@ -46,7 +47,7 @@ uv run --no-dev python -c "import torch; ok=torch.cuda.is_available(); count=tor
[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/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))" 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 == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'film_fourier_inr'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
""" """
[job] [job]
@ -79,10 +80,6 @@ required = [
"artifact_manifest.json", "artifact_manifest.json",
"checksums.txt", "checksums.txt",
"verification_report.json", "verification_report.json",
"streaming_events.jsonl",
"streaming_state.json",
"streaming_summary.json",
"processed_upload_manifest.json",
] ]
[cleanup] [cleanup]

View file

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_deeponet_branch_trunk"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/deeponet_branch_trunk.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'deeponet_branch_trunk'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_film_fourier_inr"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/film_fourier_inr.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'film_fourier_inr'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_meshgraphnet_or_point_transformer_local"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/meshgraphnet_or_point_transformer_local.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'meshgraphnet_or_point_transformer_local'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_nerf_cfd_multires"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/nerf_cfd_multires.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'nerf_cfd_multires'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_point_context_perceiver"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/point_context_perceiver.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'point_context_perceiver'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_raster_fno_unet"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/raster_fno_unet.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'raster_fno_unet'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -0,0 +1,88 @@
[run]
name = "full_70gb_model_zoo_20260726_siren_conditioned_inr"
timeout_minutes = 1440
local_artifact_dir = "artifacts/remote_runs"
max_attempts = 5
artifact_sync_interval_seconds = 60
[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",
"/data/airfrans_processed_full_50cases.tar.gz",
"/.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/model_zoo_full_70gb/siren_conditioned_inr.toml'); assert c.data.source == 'huggingface'; assert c.data.hf_repo_id == 'zacheryasc/airfrans-processed'; assert c.data.hf_path_prefix == 'processed/full'; assert c.data.cache_dir is not None; assert c.data.train_cases == 900; assert c.data.val_cases == 50; assert c.data.test_cases == 50; assert c.data.batch_size == 4096; assert c.optim.steps == 5000; assert c.model.type == 'siren_conditioned_inr'; assert c.huggingface.enabled; print('data_source=' + c.data.source + ' hf_repo=' + str(c.data.hf_repo_id) + ' hf_path=' + c.data.hf_path_prefix + ' cache_dir=' + str(c.data.cache_dir))"
"""
[job]
command = """
uv run --no-dev remote-run smoke-train configs/model_zoo_full_70gb/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

@ -35,14 +35,23 @@ def build_parser() -> argparse.ArgumentParser:
prepare_public = subparsers.add_parser( prepare_public = subparsers.add_parser(
"prepare-public-hf", "prepare-public-hf",
help="download public AirfRANS OF_dataset.zip, process it, and publish processed .npz files to HF", help="range-stream public AirfRANS, process bounded chunks, and publish verified .npz files to HF",
) )
prepare_public.add_argument("--repo-id", required=True) prepare_public.add_argument("--repo-id", default="zacheryasc/airfrans-processed")
prepare_public.add_argument("--path-in-repo", default="processed/full") prepare_public.add_argument("--path-in-repo", default="processed/full")
prepare_public.add_argument("--work-dir", default="artifacts/public_airfrans") 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("--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("--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("--min-cases", type=int, default=1000)
prepare_public.add_argument("--chunk-max-bytes", type=int, default=10 * 1024**3)
prepare_public.add_argument("--state-path")
prepare_public.add_argument("--train-cases", type=int, default=900)
prepare_public.add_argument("--val-cases", type=int, default=50)
prepare_public.add_argument("--test-cases", type=int, default=50)
prepare_public.add_argument("--split-seed", type=int, default=20260726)
prepare_public.add_argument("--limit-cases", type=int)
prepare_public.add_argument("--verify-download-limit-bytes", type=int, default=64 * 1024 * 1024)
prepare_public.add_argument("--workers", type=int, default=4)
prepare_public.add_argument("--private", action="store_true") prepare_public.add_argument("--private", action="store_true")
prepare_public.add_argument("--force", action="store_true") prepare_public.add_argument("--force", action="store_true")
prepare_public.set_defaults(command="prepare-public-hf") prepare_public.set_defaults(command="prepare-public-hf")
@ -131,18 +140,39 @@ def main(argv: list[str] | None = None) -> int:
if args.min_cases <= 0: if args.min_cases <= 0:
print("error: --min-cases must be positive", file=sys.stderr) print("error: --min-cases must be positive", file=sys.stderr)
return 1 return 1
from airfrans_frontier.raw.public import ensure_public_airfrans_processed_hf if args.chunk_max_bytes <= 0:
print("error: --chunk-max-bytes must be positive", file=sys.stderr)
return 1
if args.workers <= 0:
print("error: --workers must be positive", file=sys.stderr)
return 1
if args.limit_cases is not None and args.limit_cases <= 0:
print("error: --limit-cases must be positive", file=sys.stderr)
return 1
if args.train_cases + args.val_cases + args.test_cases <= 0:
print("error: at least one split case is required", file=sys.stderr)
return 1
from airfrans_frontier.raw.bounded_public import prepare_public_airfrans_processed_hf_bounded
try: try:
report = ensure_public_airfrans_processed_hf( report = prepare_public_airfrans_processed_hf_bounded(
repo_id=args.repo_id, repo_id=args.repo_id,
path_in_repo=args.path_in_repo, path_in_repo=args.path_in_repo,
work_dir=resolve_path(args.work_dir), work_dir=resolve_path(args.work_dir),
output_dir=resolve_path(args.output_dir), output_dir=resolve_path(args.output_dir),
source_url=args.source_url, source_url=args.source_url,
min_cases=args.min_cases, min_cases=args.min_cases,
chunk_max_bytes=args.chunk_max_bytes,
state_path=resolve_path(args.state_path) if args.state_path else None,
train_cases=args.train_cases,
val_cases=args.val_cases,
test_cases=args.test_cases,
split_seed=args.split_seed,
private=args.private, private=args.private,
force=args.force, force=args.force,
limit_cases=args.limit_cases,
verify_download_limit_bytes=args.verify_download_limit_bytes,
workers=args.workers,
) )
except (FileNotFoundError, NotADirectoryError, ValueError, RuntimeError) as exc: except (FileNotFoundError, NotADirectoryError, ValueError, RuntimeError) as exc:
print(f"error: {exc}", file=sys.stderr) print(f"error: {exc}", file=sys.stderr)

View file

@ -0,0 +1,843 @@
from __future__ import annotations
import hashlib
import json
import os
import shutil
import tempfile
import time
from concurrent.futures import ProcessPoolExecutor, as_completed
from dataclasses import dataclass
from pathlib import Path, PurePosixPath
from typing import Any, Mapping
from airfrans_frontier.raw.process import FEATURE_NAMES, TARGET_NAMES, process_raw_case_to_npz
from airfrans_frontier.raw.public import (
DEFAULT_PUBLIC_OUTPUT_DIR,
DEFAULT_PUBLIC_WORK_DIR,
PUBLIC_OF_DATASET_URL,
RangeReader,
RemoteZipMember,
_extract_remote_case_members,
_range_reader_for,
_read_zip_central_directory,
_remote_archive_case_members,
)
from airfrans_frontier.training.data import create_case_split
DEFAULT_PROCESSED_HF_REPO_ID = "zacheryasc/airfrans-processed"
DEFAULT_PROCESSED_HF_PATH = "processed/full"
DEFAULT_CHUNK_MAX_BYTES = 10 * 1024**3
DEFAULT_VERIFY_DOWNLOAD_LIMIT_BYTES = 64 * 1024 * 1024
DEFAULT_WORKERS = 4
DEFAULT_SPLIT_SEED = 20260726
DEFAULT_TRAIN_CASES = 900
DEFAULT_VAL_CASES = 50
DEFAULT_TEST_CASES = 50
_STATE_SCHEMA_VERSION = 1
@dataclass(frozen=True)
class BoundedCasePlan:
case_id: str
members: list[tuple[RemoteZipMember, PurePosixPath]]
raw_uncompressed_bytes: int
compressed_bytes: int
estimated_local_bytes: int
def prepare_public_airfrans_processed_hf_bounded(
*,
repo_id: str = DEFAULT_PROCESSED_HF_REPO_ID,
path_in_repo: str = DEFAULT_PROCESSED_HF_PATH,
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,
chunk_max_bytes: int = DEFAULT_CHUNK_MAX_BYTES,
state_path: str | Path | None = None,
train_cases: int = DEFAULT_TRAIN_CASES,
val_cases: int = DEFAULT_VAL_CASES,
test_cases: int = DEFAULT_TEST_CASES,
split_seed: int = DEFAULT_SPLIT_SEED,
private: bool = False,
force: bool = False,
limit_cases: int | None = None,
verify_download_limit_bytes: int = DEFAULT_VERIFY_DOWNLOAD_LIMIT_BYTES,
workers: int = DEFAULT_WORKERS,
) -> dict[str, Any]:
"""Range-stream, process, upload, verify, and discard public AirfRANS chunks.
The raw public ZIP is read only through RangeReader. At no point does this
routine download the full archive or extract the full raw tree. Processed
.npz files are staged only until their chunk commit has been verified in HF.
"""
if min_cases <= 0:
raise ValueError("min_cases must be positive")
if chunk_max_bytes <= 0:
raise ValueError("chunk_max_bytes must be positive")
if workers <= 0:
raise ValueError("workers must be positive")
if limit_cases is not None and limit_cases <= 0:
raise ValueError("limit_cases must be positive when provided")
requested_split_cases = train_cases + val_cases + test_cases
if requested_split_cases <= 0:
raise ValueError("at least one split case is required")
started = time.time()
prefix = path_in_repo.strip("/")
work_root = Path(work_dir).expanduser()
staging_root = Path(output_dir).expanduser()
scratch_root = work_root / "bounded_raw_scratch"
manifest_root = work_root / "chunk_manifests"
state_file = Path(state_path).expanduser() if state_path is not None else work_root / "bounded_prepare_state.json"
if force:
_remove_file_best_effort(state_file)
_remove_tree_best_effort(scratch_root)
_remove_tree_best_effort(staging_root)
work_root.mkdir(parents=True, exist_ok=True)
staging_root.mkdir(parents=True, exist_ok=True)
manifest_root.mkdir(parents=True, exist_ok=True)
_remove_tree_best_effort(scratch_root)
scratch_root.mkdir(parents=True, exist_ok=True)
token = _required_secret("HF_TOKEN", "Hugging Face processed dataset publishing")
api = _create_dataset_repo(repo_id=repo_id, private=private, token=token)
reader = _range_reader_for(source_url)
members = _read_zip_central_directory(reader)
case_members = _remote_archive_case_members(members)
all_case_ids = sorted(case_members)
if len(all_case_ids) < min_cases:
raise RuntimeError(f"AirfRANS archive has {len(all_case_ids)} cases; expected at least {min_cases}")
selected_case_ids = all_case_ids[:limit_cases] if limit_cases is not None else all_case_ids
if len(selected_case_ids) < min_cases:
raise RuntimeError(
f"Selected {len(selected_case_ids)} cases but min_cases={min_cases}; lower min_cases or raise limit_cases"
)
if requested_split_cases > len(selected_case_ids):
raise RuntimeError(
f"Requested {requested_split_cases} split cases but only {len(selected_case_ids)} selected cases are available"
)
plans = [_case_plan(case_id, case_members[case_id]) for case_id in selected_case_ids]
state = _load_state(state_file)
state = _compatible_or_fresh_state(
state,
repo_id=repo_id,
path_in_repo=prefix,
source_url=source_url,
chunk_max_bytes=chunk_max_bytes,
min_cases=min_cases,
selected_case_count=len(selected_case_ids),
train_cases=train_cases,
val_cases=val_cases,
test_cases=test_cases,
split_seed=split_seed,
source_bytes=reader.size,
)
state.setdefault("phase", "processing")
state.setdefault("uploaded_cases", {})
state.setdefault("pending_records", [])
state.setdefault("chunks", [])
state.setdefault("next_chunk_index", 1)
state["workers"] = workers
_write_state(state_file, state)
pending_records = _existing_pending_records(state.get("pending_records", []))
if len(pending_records) != len(state.get("pending_records", [])):
state["pending_records"] = pending_records
_write_state(state_file, state)
_remove_stale_staged_npz(staging_root, pending_records)
uploaded_cases = state["uploaded_cases"]
pending_case_ids = {str(record["case_id"]) for record in pending_records}
pending_chunk_bytes = _pending_chunk_bytes(pending_records)
processed_now = 0
plan_index = 0
while plan_index < len(plans):
while plan_index < len(plans) and (
plans[plan_index].case_id in uploaded_cases or plans[plan_index].case_id in pending_case_ids
):
plan_index += 1
if plan_index >= len(plans):
break
next_plan = plans[plan_index]
if pending_records and pending_chunk_bytes + next_plan.estimated_local_bytes > chunk_max_bytes:
_upload_pending_chunk(
api=api,
repo_id=repo_id,
prefix=prefix,
token=token,
state=state,
state_file=state_file,
manifest_root=manifest_root,
pending_records=pending_records,
verify_download_limit_bytes=verify_download_limit_bytes,
)
pending_records = []
pending_case_ids = set()
pending_chunk_bytes = 0
continue
chunk_plans: list[BoundedCasePlan] = []
planned_chunk_bytes = pending_chunk_bytes
while plan_index < len(plans):
plan = plans[plan_index]
plan_index += 1
if plan.case_id in uploaded_cases or plan.case_id in pending_case_ids:
continue
if chunk_plans and planned_chunk_bytes + plan.estimated_local_bytes > chunk_max_bytes:
plan_index -= 1
break
if pending_records and not chunk_plans and planned_chunk_bytes + plan.estimated_local_bytes > chunk_max_bytes:
plan_index -= 1
break
chunk_plans.append(plan)
planned_chunk_bytes += plan.estimated_local_bytes
if planned_chunk_bytes >= chunk_max_bytes:
break
if not chunk_plans:
continue
for record in _process_case_plans(
source_url=source_url,
plans=chunk_plans,
scratch_root=scratch_root,
staging_root=staging_root,
prefix=prefix,
workers=workers,
):
pending_records.append(record)
pending_case_ids.add(str(record["case_id"]))
pending_chunk_bytes += _record_chunk_bytes(record)
state["pending_records"] = pending_records
state["phase"] = "processing"
state["ranged_bytes_read"] = int(state.get("ranged_bytes_read", reader.bytes_read) or 0) + int(
record.get("ranged_bytes_read", 0) or 0
)
state["updated_at"] = time.time()
_write_state(state_file, state)
processed_now += 1
if pending_chunk_bytes >= chunk_max_bytes:
_upload_pending_chunk(
api=api,
repo_id=repo_id,
prefix=prefix,
token=token,
state=state,
state_file=state_file,
manifest_root=manifest_root,
pending_records=pending_records,
verify_download_limit_bytes=verify_download_limit_bytes,
)
pending_records = []
pending_case_ids = set()
pending_chunk_bytes = 0
if pending_records:
_upload_pending_chunk(
api=api,
repo_id=repo_id,
prefix=prefix,
token=token,
state=state,
state_file=state_file,
manifest_root=manifest_root,
pending_records=pending_records,
verify_download_limit_bytes=verify_download_limit_bytes,
)
pending_records = []
uploaded_records = sorted(state["uploaded_cases"].values(), key=lambda item: str(item["case_id"]))
if len(uploaded_records) < min_cases:
raise RuntimeError(f"Uploaded {len(uploaded_records)} cases but expected at least {min_cases}")
_verify_uploaded_records(
api=api,
repo_id=repo_id,
records=uploaded_records,
token=token,
verify_download_limit_bytes=verify_download_limit_bytes,
)
final_manifest_path = _write_final_manifest(
manifest_root=manifest_root,
repo_id=repo_id,
prefix=prefix,
source_url=source_url,
source_bytes=reader.size,
records=uploaded_records,
chunks=state.get("chunks", []),
train_cases=train_cases,
val_cases=val_cases,
test_cases=test_cases,
split_seed=split_seed,
chunk_max_bytes=chunk_max_bytes,
elapsed_seconds=time.time() - started,
)
final_manifest_record = _file_upload_record(
local_path=final_manifest_path,
repo_path=f"{prefix}/hf_dataset_manifest.json" if prefix else "hf_dataset_manifest.json",
kind="final_manifest",
)
commit = _commit_files(
api=api,
repo_id=repo_id,
records=[final_manifest_record],
message="Publish bounded AirfRANS processed dataset manifest",
)
_verify_uploaded_records(
api=api,
repo_id=repo_id,
records=[final_manifest_record],
token=token,
verify_download_limit_bytes=verify_download_limit_bytes,
)
repo_status = _dataset_status(api=api, repo_id=repo_id, prefix=prefix)
if repo_status["npz_file_count"] < min_cases:
raise RuntimeError(f"HF dataset has {repo_status['npz_file_count']} .npz files under {prefix}; expected at least {min_cases}")
if not repo_status["has_manifest"]:
raise RuntimeError(f"HF dataset is missing hf_dataset_manifest.json under {prefix}")
state.update(
{
"phase": "published",
"pending_records": [],
"final_manifest": {**final_manifest_record, "commit": _commit_payload(commit)},
"repo_status": repo_status,
"ranged_bytes_read": int(state.get("ranged_bytes_read", reader.bytes_read) or 0),
"finished_at": time.time(),
"updated_at": time.time(),
}
)
_write_state(state_file, state)
_remove_tree_best_effort(scratch_root)
_remove_uploaded_staging_files(uploaded_records)
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_bounded_chunks",
"source_bytes": reader.size,
"ranged_bytes_read": int(state.get("ranged_bytes_read", reader.bytes_read) or 0),
"chunk_max_bytes": chunk_max_bytes,
"workers": workers,
"state_path": str(state_file),
"final_manifest_path": str(final_manifest_path),
"processed_case_count": len(uploaded_records),
"processed_now": processed_now,
"chunks_uploaded": len(state.get("chunks", [])),
"elapsed_seconds": time.time() - started,
**repo_status,
}
def _case_plan(case_id: str, members: list[tuple[RemoteZipMember, PurePosixPath]]) -> BoundedCasePlan:
raw_uncompressed = sum(member.file_size for member, _relative in members if not member.is_dir)
compressed = sum(member.compress_size for member, _relative in members if not member.is_dir)
return BoundedCasePlan(
case_id=case_id,
members=members,
raw_uncompressed_bytes=raw_uncompressed,
compressed_bytes=compressed,
estimated_local_bytes=max(raw_uncompressed, compressed),
)
def _process_case_to_staging(
*,
reader: RangeReader,
plan: BoundedCasePlan,
scratch_root: Path,
staging_root: Path,
prefix: str,
) -> dict[str, Any]:
case_dir = scratch_root / plan.case_id
target_path = staging_root / f"{plan.case_id}.npz"
try:
_extract_remote_case_members(reader, plan.members, scratch_root)
processing_record, points = process_raw_case_to_npz(case_dir, staging_root, force=True)
finally:
_remove_tree_best_effort(case_dir)
if not target_path.is_file():
raise RuntimeError(f"Processed AirfRANS case missing staged .npz: {target_path}")
record = _file_upload_record(
local_path=target_path,
repo_path=f"{prefix}/{target_path.name}" if prefix else target_path.name,
kind="case_npz",
)
record.update(
{
"case_id": plan.case_id,
"points": int(points),
"raw_uncompressed_bytes": plan.raw_uncompressed_bytes,
"compressed_bytes": plan.compressed_bytes,
"estimated_local_bytes": plan.estimated_local_bytes,
"metadata": processing_record.get("metadata"),
"processed_at": time.time(),
}
)
return record
def _process_case_plans(
*,
source_url: str,
plans: list[BoundedCasePlan],
scratch_root: Path,
staging_root: Path,
prefix: str,
workers: int,
):
if workers == 1 or len(plans) == 1:
reader = _range_reader_for(source_url)
for plan in plans:
before = reader.bytes_read
record = _process_case_to_staging(
reader=reader,
plan=plan,
scratch_root=scratch_root,
staging_root=staging_root,
prefix=prefix,
)
record["ranged_bytes_read"] = reader.bytes_read - before
yield record
return
with ProcessPoolExecutor(max_workers=min(workers, len(plans))) as executor:
futures = [
executor.submit(
_process_case_plan_worker,
source_url,
plan,
str(scratch_root),
str(staging_root),
prefix,
)
for plan in plans
]
for future in as_completed(futures):
yield future.result()
def _process_case_plan_worker(
source_url: str,
plan: BoundedCasePlan,
scratch_root: str,
staging_root: str,
prefix: str,
) -> dict[str, Any]:
reader = _range_reader_for(source_url)
before = reader.bytes_read
record = _process_case_to_staging(
reader=reader,
plan=plan,
scratch_root=Path(scratch_root),
staging_root=Path(staging_root),
prefix=prefix,
)
record["ranged_bytes_read"] = reader.bytes_read - before
return record
def _upload_pending_chunk(
*,
api: Any,
repo_id: str,
prefix: str,
token: str,
state: dict[str, Any],
state_file: Path,
manifest_root: Path,
pending_records: list[dict[str, Any]],
verify_download_limit_bytes: int,
) -> None:
if not pending_records:
return
chunk_index = int(state.get("next_chunk_index", 1) or 1)
uploaded_at = time.time()
chunk_manifest_path = manifest_root / f"chunk_{chunk_index:04d}.json"
chunk_payload = {
"schema_version": 1,
"chunk_index": chunk_index,
"case_count": len(pending_records),
"case_ids": [record["case_id"] for record in pending_records],
"total_bytes": sum(int(record["bytes"]) for record in pending_records),
"records": [_public_record(record) for record in pending_records],
"created_at": uploaded_at,
}
_atomic_write_json(chunk_manifest_path, chunk_payload)
manifest_record = _file_upload_record(
local_path=chunk_manifest_path,
repo_path=f"{prefix}/manifests/chunks/chunk_{chunk_index:04d}.json" if prefix else f"manifests/chunks/chunk_{chunk_index:04d}.json",
kind="chunk_manifest",
)
records_to_upload = [*pending_records, manifest_record]
commit = _commit_files(
api=api,
repo_id=repo_id,
records=records_to_upload,
message=f"Upload bounded AirfRANS processed chunk {chunk_index:04d}",
)
_verify_uploaded_records(
api=api,
repo_id=repo_id,
records=records_to_upload,
token=token,
verify_download_limit_bytes=verify_download_limit_bytes,
)
uploaded_cases = state.setdefault("uploaded_cases", {})
for record in pending_records:
uploaded_cases[str(record["case_id"])] = {**record, "uploaded_at": uploaded_at}
state.setdefault("chunks", []).append(
{
"chunk_index": chunk_index,
"case_count": len(pending_records),
"case_ids": [record["case_id"] for record in pending_records],
"total_bytes": sum(int(record["bytes"]) for record in pending_records),
"manifest": {**manifest_record, "uploaded_at": uploaded_at},
"commit": _commit_payload(commit),
"uploaded_at": uploaded_at,
}
)
state["pending_records"] = []
state["next_chunk_index"] = chunk_index + 1
state["phase"] = "uploaded_chunk"
state["updated_at"] = time.time()
_write_state(state_file, state)
_remove_uploaded_staging_files(pending_records)
def _commit_files(*, api: Any, repo_id: str, records: list[dict[str, Any]], message: str) -> Any:
try:
from huggingface_hub import CommitOperationAdd
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required for AirfRANS processed dataset publishing") from exc
operations = [
CommitOperationAdd(path_in_repo=str(record["repo_path"]), path_or_fileobj=str(record["local_path"]))
for record in records
]
return api.create_commit(repo_id=repo_id, repo_type="dataset", operations=operations, commit_message=message)
def _write_final_manifest(
*,
manifest_root: Path,
repo_id: str,
prefix: str,
source_url: str,
source_bytes: int,
records: list[dict[str, Any]],
chunks: list[dict[str, Any]],
train_cases: int,
val_cases: int,
test_cases: int,
split_seed: int,
chunk_max_bytes: int,
elapsed_seconds: float,
) -> Path:
case_ids = [str(record["case_id"]) for record in records]
split = create_case_split(
tuple(case_ids),
train_cases=train_cases,
val_cases=val_cases,
test_cases=test_cases,
seed=split_seed,
)
payload = {
"schema_version": 2,
"repo_id": repo_id,
"repo_type": "dataset",
"repo_url": f"https://huggingface.co/datasets/{repo_id}",
"path_in_repo": prefix,
"source": "public_zip_range_bounded_chunks",
"source_url": source_url,
"source_bytes": source_bytes,
"chunk_max_bytes": chunk_max_bytes,
"case_count": len(records),
"npz_file_count": len(records),
"total_bytes": sum(int(record["bytes"]) for record in records),
"total_points": sum(int(record.get("points", 0) or 0) for record in records),
"feature_names": FEATURE_NAMES.tolist(),
"target_names": TARGET_NAMES.tolist(),
"cases": [_public_record(record) for record in records],
"chunks": chunks,
"split_compatibility": {
"seed": split_seed,
"train_cases": train_cases,
"val_cases": val_cases,
"test_cases": test_cases,
"train_ids": list(split.train_ids),
"val_ids": list(split.val_ids),
"test_ids": list(split.test_ids),
},
"created_at": time.time(),
"elapsed_seconds": elapsed_seconds,
}
path = manifest_root / "hf_dataset_manifest.json"
_atomic_write_json(path, payload)
return path
def _file_upload_record(*, local_path: Path, repo_path: str, kind: str) -> dict[str, Any]:
return {
"kind": kind,
"local_path": str(local_path),
"repo_path": repo_path,
"bytes": local_path.stat().st_size,
"sha256": _sha256_file(local_path),
}
def _public_record(record: Mapping[str, Any]) -> dict[str, Any]:
return {
key: value
for key, value in record.items()
if key not in {"local_path"} and value is not None
}
def _verify_uploaded_records(
*,
api: Any,
repo_id: str,
records: list[Mapping[str, Any]],
token: str,
verify_download_limit_bytes: int,
) -> None:
metadata = _repo_file_metadata(api=api, repo_id=repo_id)
missing: list[str] = []
mismatches: list[str] = []
for record in records:
repo_path = str(record["repo_path"])
info = metadata.get(repo_path)
if info is None:
missing.append(repo_path)
continue
expected_size = int(record["bytes"])
remote_size = _metadata_size(info)
if remote_size is not None and remote_size != expected_size:
mismatches.append(f"{repo_path}: size {remote_size} != {expected_size}")
continue
expected_sha = str(record["sha256"])
remote_sha = _metadata_sha256(info)
if remote_sha is None:
remote_sha = _downloaded_sha256(
repo_id=repo_id,
repo_path=repo_path,
token=token,
size=expected_size,
limit=verify_download_limit_bytes,
)
if remote_sha != expected_sha:
mismatches.append(f"{repo_path}: sha256 {remote_sha} != {expected_sha}")
if missing or mismatches:
details = []
if missing:
details.append("missing=" + ",".join(missing))
if mismatches:
details.append("mismatches=" + "; ".join(mismatches))
raise RuntimeError("HF upload verification failed: " + " ".join(details))
def _repo_file_metadata(*, api: Any, repo_id: str) -> dict[str, Any]:
try:
info = api.repo_info(repo_id=repo_id, repo_type="dataset", files_metadata=True)
siblings = getattr(info, "siblings", []) or []
return {str(getattr(item, "rfilename", "")): item for item in siblings if getattr(item, "rfilename", None)}
except Exception:
files = api.list_repo_files(repo_id=repo_id, repo_type="dataset") if hasattr(api, "list_repo_files") else []
return {str(path): {"rfilename": path} for path in files}
def _dataset_status(*, api: Any, repo_id: str, prefix: str) -> dict[str, Any]:
files = api.list_repo_files(repo_id=repo_id, repo_type="dataset") if hasattr(api, "list_repo_files") else list(_repo_file_metadata(api=api, repo_id=repo_id))
base = f"{prefix}/" if prefix else ""
npz_count = sum(1 for path in files if str(path).startswith(base) and str(path).endswith(".npz"))
has_manifest = any(str(path) == f"{base}hf_dataset_manifest.json" for path in files)
return {"file_count": len(files), "npz_file_count": npz_count, "has_manifest": has_manifest}
def _metadata_size(info: Any) -> int | None:
if isinstance(info, Mapping):
size = info.get("size")
else:
size = getattr(info, "size", None)
return int(size) if size is not None else None
def _metadata_sha256(info: Any) -> str | None:
lfs = info.get("lfs") if isinstance(info, Mapping) else getattr(info, "lfs", None)
if isinstance(lfs, Mapping):
value = lfs.get("sha256") or lfs.get("oid")
return str(value) if value else None
value = getattr(lfs, "sha256", None) or getattr(lfs, "oid", None)
return str(value) if value else None
def _downloaded_sha256(*, repo_id: str, repo_path: str, token: str, size: int, limit: int) -> str:
if size > limit:
raise RuntimeError(f"remote metadata for {repo_path} has no LFS sha256 and size {size} exceeds verify limit {limit}")
try:
from huggingface_hub import hf_hub_download
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required for AirfRANS processed dataset verification") from exc
with tempfile.TemporaryDirectory() as tmp:
downloaded = hf_hub_download(
repo_id=repo_id,
repo_type="dataset",
filename=repo_path,
token=token,
local_dir=tmp,
)
return _sha256_file(Path(downloaded))
def _create_dataset_repo(*, repo_id: str, private: bool, token: str) -> Any:
try:
from huggingface_hub import HfApi
except ModuleNotFoundError as exc:
raise RuntimeError("huggingface_hub is required for AirfRANS processed dataset publishing") from exc
api = HfApi(token=token)
api.create_repo(repo_id=repo_id, repo_type="dataset", private=private, exist_ok=True)
return api
def _compatible_or_fresh_state(
state: dict[str, Any],
*,
repo_id: str,
path_in_repo: str,
source_url: str,
chunk_max_bytes: int,
min_cases: int,
selected_case_count: int,
train_cases: int,
val_cases: int,
test_cases: int,
split_seed: int,
source_bytes: int,
) -> dict[str, Any]:
expected = {
"schema_version": _STATE_SCHEMA_VERSION,
"repo_id": repo_id,
"path_in_repo": path_in_repo,
"source_url": source_url,
"chunk_max_bytes": chunk_max_bytes,
"min_cases": min_cases,
"selected_case_count": selected_case_count,
"train_cases": train_cases,
"val_cases": val_cases,
"test_cases": test_cases,
"split_seed": split_seed,
"source_bytes": source_bytes,
}
if not state:
return {**expected, "created_at": time.time(), "updated_at": time.time()}
for key, value in expected.items():
if state.get(key) != value:
raise RuntimeError(
f"Existing bounded AirfRANS state is for a different run: {key}={state.get(key)!r} expected {value!r}; pass force=True to restart"
)
return state
def _load_state(path: Path) -> dict[str, Any]:
if not path.is_file():
return {}
try:
data = json.loads(path.read_text())
except json.JSONDecodeError as exc:
raise RuntimeError(f"Invalid bounded AirfRANS state JSON: {path}: {exc}") from exc
if not isinstance(data, dict):
raise RuntimeError(f"Invalid bounded AirfRANS state JSON root: {path}")
return data
def _existing_pending_records(records: Any) -> list[dict[str, Any]]:
if not isinstance(records, list):
return []
existing: list[dict[str, Any]] = []
for record in records:
if not isinstance(record, dict):
continue
path = Path(str(record.get("local_path", "")))
if path.is_file():
existing.append(record)
return existing
def _remove_stale_staged_npz(staging_root: Path, pending_records: list[Mapping[str, Any]]) -> None:
pending_paths = {Path(str(record.get("local_path", ""))).resolve() for record in pending_records}
for path in staging_root.glob("*.npz"):
if path.resolve() not in pending_paths:
_remove_file_best_effort(path)
def _pending_chunk_bytes(records: list[Mapping[str, Any]]) -> int:
return sum(_record_chunk_bytes(record) for record in records)
def _record_chunk_bytes(record: Mapping[str, Any]) -> int:
return max(int(record.get("estimated_local_bytes", 0) or 0), int(record.get("bytes", 0) or 0))
def _remove_uploaded_staging_files(records: list[Mapping[str, Any]]) -> None:
for record in records:
path = Path(str(record.get("local_path", "")))
_remove_file_best_effort(path)
def _required_secret(name: str, purpose: str) -> str:
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
raise RuntimeError(f"{name} env var or local secret file is required for {purpose}")
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),
}
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _atomic_write_json(path: Path, payload: Mapping[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_name(f"{path.name}.tmp")
tmp.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
tmp.replace(path)
def _write_state(path: Path, state: Mapping[str, Any]) -> None:
_atomic_write_json(path, state)
def _remove_tree_best_effort(path: Path) -> None:
if path.exists():
shutil.rmtree(path, ignore_errors=True)
def _remove_file_best_effort(path: Path) -> None:
try:
path.unlink()
except FileNotFoundError:
return

View file

@ -4,13 +4,14 @@ import json
import os import os
import shutil import shutil
import struct import struct
import http.client
import time import time
import urllib.error import urllib.error
import urllib.request import urllib.request
import urllib.parse import urllib.parse
import zlib import zlib
import zipfile import zipfile
from dataclasses import dataclass from dataclasses import dataclass, replace
from pathlib import Path, PurePosixPath from pathlib import Path, PurePosixPath
from typing import Any, Protocol from typing import Any, Protocol
@ -28,6 +29,10 @@ _LOCAL_FILE_HEADER_SIGNATURE = 0x04034B50
_ZIP64_EXTRA_ID = 0x0001 _ZIP64_EXTRA_ID = 0x0001
_ZIP64_LIMIT_16 = 0xFFFF _ZIP64_LIMIT_16 = 0xFFFF
_ZIP64_LIMIT_32 = 0xFFFFFFFF _ZIP64_LIMIT_32 = 0xFFFFFFFF
_HTTP_RANGE_READ_TIMEOUT_SECONDS = 60
_HTTP_RANGE_READ_MAX_ATTEMPTS = 4
_HTTP_RANGE_READ_RETRY_BASE_SECONDS = 2.0
class RangeReader(Protocol): class RangeReader(Protocol):
@ -45,6 +50,7 @@ class RemoteZipMember:
compress_size: int compress_size: int
file_size: int file_size: int
header_offset: int header_offset: int
next_header_offset: int | None = None
@property @property
def is_dir(self) -> bool: def is_dir(self) -> bool:
@ -91,23 +97,41 @@ class HttpRangeReader:
if length == 0: if length == 0:
return b"" return b""
end = start + length - 1 end = start + length - 1
last_error: RuntimeError | None = None
for attempt in range(1, _HTTP_RANGE_READ_MAX_ATTEMPTS + 1):
request = urllib.request.Request(self.url, headers={"Range": f"bytes={start}-{end}"}) request = urllib.request.Request(self.url, headers={"Range": f"bytes={start}-{end}"})
try: try:
with urllib.request.urlopen(request, timeout=60) as response: with urllib.request.urlopen(request, timeout=_HTTP_RANGE_READ_TIMEOUT_SECONDS) as response:
status = getattr(response, "status", None) status = getattr(response, "status", None)
data = response.read() data = response.read()
except urllib.error.HTTPError as exc: except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace") 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 last_error = RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc.code} {body}")
except OSError as exc: if exc.code in {408, 429} or 500 <= exc.code < 600:
raise RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc}") from exc if attempt < _HTTP_RANGE_READ_MAX_ATTEMPTS:
time.sleep(_HTTP_RANGE_READ_RETRY_BASE_SECONDS * attempt)
continue
raise last_error from exc
except (http.client.IncompleteRead, OSError) as exc:
last_error = RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}: {exc}")
if attempt < _HTTP_RANGE_READ_MAX_ATTEMPTS:
time.sleep(_HTTP_RANGE_READ_RETRY_BASE_SECONDS * attempt)
continue
raise last_error from exc
if status != 206: if status != 206:
raise RuntimeError(f"Server did not honor HTTP Range for {self.url}: status={status}") raise RuntimeError(f"Server did not honor HTTP Range for {self.url}: status={status}")
if len(data) != length: if len(data) != length:
raise RuntimeError(f"HTTP range read returned {len(data)} bytes; expected {length}") last_error = RuntimeError(f"HTTP range read returned {len(data)} bytes; expected {length}")
if attempt < _HTTP_RANGE_READ_MAX_ATTEMPTS:
time.sleep(_HTTP_RANGE_READ_RETRY_BASE_SECONDS * attempt)
continue
raise last_error
self.bytes_read += len(data) self.bytes_read += len(data)
return data return data
raise last_error or RuntimeError(f"HTTP range read failed for {self.url} bytes={start}-{end}")
def ensure_public_airfrans_processed_hf( def ensure_public_airfrans_processed_hf(
*, *,
repo_id: str, repo_id: str,
@ -317,7 +341,7 @@ def _read_zip_central_directory(reader: RangeReader) -> list[RemoteZipMember]:
): ):
total_entries, central_size, central_offset = _read_zip64_central_directory_locator(reader, eocd_offset) total_entries, central_size, central_offset = _read_zip64_central_directory_locator(reader, eocd_offset)
central = reader.read_range(central_offset, central_size) central = reader.read_range(central_offset, central_size)
members = _parse_central_directory(central, expected_entries=total_entries) members = _annotate_next_header_offsets(_parse_central_directory(central, expected_entries=total_entries), central_offset=central_offset)
print(f"range_stream_airfrans_zip_members={len(members)}", flush=True) print(f"range_stream_airfrans_zip_members={len(members)}", flush=True)
return members return members
@ -398,6 +422,16 @@ def _parse_central_directory(central: bytes, *, expected_entries: int) -> list[R
return members return members
def _annotate_next_header_offsets(members: list[RemoteZipMember], *, central_offset: int) -> list[RemoteZipMember]:
next_offsets: dict[int, int] = {}
ordered = sorted(enumerate(members), key=lambda item: item[1].header_offset)
for position, (original_index, _member) in enumerate(ordered):
next_offsets[original_index] = (
ordered[position + 1][1].header_offset if position + 1 < len(ordered) else central_offset
)
return [replace(member, next_header_offset=next_offsets[index]) for index, member in enumerate(members)]
def _apply_zip64_extra(extra: bytes, *, file_size: int, compress_size: int, header_offset: int) -> tuple[int, int, int]: def _apply_zip64_extra(extra: bytes, *, file_size: int, compress_size: int, header_offset: int) -> tuple[int, int, int]:
values_needed = [ values_needed = [
file_size == _ZIP64_LIMIT_32, file_size == _ZIP64_LIMIT_32,
@ -502,7 +536,21 @@ def _extract_remote_case_members(
root: Path, root: Path,
) -> None: ) -> None:
resolved_root = root.resolve() resolved_root = root.resolve()
for member, relative in members: entries = sorted(members, key=lambda item: item[0].header_offset)
span = _contiguous_case_span([member for member, _relative in entries])
if span is not None:
span_start, span_end = span
archive_bytes = reader.read_range(span_start, span_end - span_start)
for member, relative in entries:
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)
target.write_bytes(_read_remote_member_payload_from_span(member, archive_bytes, span_start))
return
for member, relative in entries:
target = _safe_relative_target(root, relative, resolved_root=resolved_root) target = _safe_relative_target(root, relative, resolved_root=resolved_root)
if member.is_dir: if member.is_dir:
target.mkdir(parents=True, exist_ok=True) target.mkdir(parents=True, exist_ok=True)
@ -512,9 +560,20 @@ def _extract_remote_case_members(
target.write_bytes(payload) target.write_bytes(payload)
def _contiguous_case_span(members: list[RemoteZipMember]) -> tuple[int, int] | None:
if not members:
return None
ordered = sorted(members, key=lambda member: member.header_offset)
for current, following in zip(ordered, ordered[1:]):
if current.next_header_offset != following.header_offset:
return None
span_end = ordered[-1].next_header_offset
if span_end is None or span_end <= ordered[0].header_offset:
return None
return ordered[0].header_offset, span_end
def _read_remote_member_payload(reader: RangeReader, member: RemoteZipMember) -> bytes: 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) local_header = reader.read_range(member.header_offset, 30)
( (
signature, signature,
@ -533,6 +592,40 @@ def _read_remote_member_payload(reader: RangeReader, member: RemoteZipMember) ->
raise RuntimeError(f"Invalid local ZIP header for member: {member.filename}") raise RuntimeError(f"Invalid local ZIP header for member: {member.filename}")
data_offset = member.header_offset + 30 + filename_length + extra_length data_offset = member.header_offset + 30 + filename_length + extra_length
compressed = reader.read_range(data_offset, member.compress_size) compressed = reader.read_range(data_offset, member.compress_size)
return _decode_remote_member_payload(member, compressed)
def _read_remote_member_payload_from_span(member: RemoteZipMember, archive_bytes: bytes, span_start: int) -> bytes:
local_header_offset = member.header_offset - span_start
local_header = archive_bytes[local_header_offset : local_header_offset + 30]
if len(local_header) != 30:
raise RuntimeError(f"Truncated local ZIP header for member: {member.filename}")
(
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 = local_header_offset + 30 + filename_length + extra_length
data_end = data_offset + member.compress_size
compressed = archive_bytes[data_offset:data_end]
if len(compressed) != member.compress_size:
raise RuntimeError(f"Truncated ZIP member payload for member: {member.filename}")
return _decode_remote_member_payload(member, compressed)
def _decode_remote_member_payload(member: RemoteZipMember, compressed: bytes) -> bytes:
if member.flag_bits & 0x1:
raise RuntimeError(f"Encrypted ZIP member is unsupported: {member.filename}")
if member.compress_type == 0: if member.compress_type == 0:
payload = compressed payload = compressed
elif member.compress_type == 8: elif member.compress_type == 8:

View file

@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence from collections.abc import Callable, Mapping, Sequence
import re
import time import time
from typing import Any from typing import Any
@ -20,6 +21,7 @@ def reconcile_cleanup(
sky_state: Any, sky_state: Any,
vast_instances: Sequence[Mapping[str, Any]], vast_instances: Sequence[Mapping[str, Any]],
known_run_ids: Sequence[str] = (), known_run_ids: Sequence[str] = (),
ignored_label_prefixes: Sequence[str] = (),
destroy_orphans: bool = False, destroy_orphans: bool = False,
destroy_instance: Callable[[int], Any] | None = None, destroy_instance: Callable[[int], Any] | None = None,
now: float | None = None, now: float | None = None,
@ -32,17 +34,20 @@ def reconcile_cleanup(
checked_at = time.time() if now is None else float(now) checked_at = time.time() if now is None else float(now)
sky_refs = _extract_sky_refs(sky_state) sky_refs = _extract_sky_refs(sky_state)
known_runs = tuple(known_run_ids) known_runs = tuple(dict.fromkeys((*known_run_ids, *sorted(sky_refs["run_ids"]))))
ignored_prefixes = tuple(prefix for prefix in ignored_label_prefixes if prefix)
records: list[dict[str, Any]] = [] records: list[dict[str, Any]] = []
for instance in vast_instances: for instance in vast_instances:
instance_id = _instance_id(instance) instance_id = _instance_id(instance)
status = _status(instance) status = _status(instance)
label = _instance_label(instance)
associated_run_id = _associated_run_id(instance, known_runs) 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) sky_knows = _sky_knows_instance(sky_refs, instance_id=instance_id, run_id=associated_run_id)
live = _is_live_status(status) live = _is_live_status(status)
unexpected_live = bool(live and not sky_knows) ignored_external = bool(live and not sky_knows and _matches_ignored_label(label, ignored_prefixes))
action = "none" unexpected_live = bool(live and not sky_knows and not ignored_external)
result = "not_needed" action = "ignore_external" if ignored_external else "none"
result = "ignored" if ignored_external else "not_needed"
error = None error = None
if unexpected_live: if unexpected_live:
action = "destroy_orphan" if destroy_orphans else "report_orphan" action = "destroy_orphan" if destroy_orphans else "report_orphan"
@ -67,6 +72,8 @@ def reconcile_cleanup(
"gpu_type": _first_present(instance, "gpu_name", "gpu", "gpu_type"), "gpu_type": _first_present(instance, "gpu_name", "gpu", "gpu_type"),
"gpu_count": _first_present(instance, "num_gpus", "gpu_count", "gpus"), "gpu_count": _first_present(instance, "num_gpus", "gpu_count", "gpus"),
"status": status, "status": status,
"label": label,
"ignored_external": ignored_external,
"associated_run_id": associated_run_id, "associated_run_id": associated_run_id,
"hourly_cost": _first_present(instance, "dph_total", "hourly_cost", "cost_per_hour"), "hourly_cost": _first_present(instance, "dph_total", "hourly_cost", "cost_per_hour"),
"sky_known": sky_knows, "sky_known": sky_knows,
@ -83,7 +90,9 @@ def reconcile_cleanup(
"sky_instance_ids": sorted(sky_refs["instance_ids"]), "sky_instance_ids": sorted(sky_refs["instance_ids"]),
"sky_run_ids": sorted(sky_refs["run_ids"]), "sky_run_ids": sorted(sky_refs["run_ids"]),
"destroy_orphans": destroy_orphans, "destroy_orphans": destroy_orphans,
"ignored_label_prefixes": sorted(ignored_prefixes),
"unexpected_live_count": sum(1 for record in records if record["unexpected_live"]), "unexpected_live_count": sum(1 for record in records if record["unexpected_live"]),
"ignored_external_count": sum(1 for record in records if record["ignored_external"]),
"destroy_requested_count": sum(1 for record in records if record["cleanup_result"] == "destroy_requested"), "destroy_requested_count": sum(1 for record in records if record["cleanup_result"] == "destroy_requested"),
"instances": records, "instances": records,
} }
@ -99,8 +108,8 @@ def _walk_sky(value: Any, refs: dict[str, set[str]]) -> None:
if isinstance(value, Mapping): if isinstance(value, Mapping):
for key, item in value.items(): for key, item in value.items():
key_text = str(key).lower() key_text = str(key).lower()
if key_text in {"id", "instance_id", "vast_instance_id"}: if key_text in {"id", "instance_id", "vast_instance_id", "node_id", "node_ids", "node_name", "node_names"}:
_add_ref(refs["instance_ids"], item) _add_instance_refs(refs["instance_ids"], item)
elif key_text in {"name", "cluster", "cluster_name", "run_id", "label"}: elif key_text in {"name", "cluster", "cluster_name", "run_id", "label"}:
_add_ref(refs["run_ids"], item) _add_ref(refs["run_ids"], item)
_walk_sky(item, refs) _walk_sky(item, refs)
@ -109,6 +118,19 @@ def _walk_sky(value: Any, refs: dict[str, set[str]]) -> None:
_walk_sky(item, refs) _walk_sky(item, refs)
def _add_instance_refs(target: set[str], value: Any) -> None:
if isinstance(value, str):
for item in re.split(r"[^0-9]+", value):
if item:
target.add(item)
return
if isinstance(value, (list, tuple, set)):
for item in value:
_add_instance_refs(target, item)
return
_add_ref(target, value)
def _add_ref(target: set[str], value: Any) -> None: def _add_ref(target: set[str], value: Any) -> None:
if isinstance(value, bool) or value is None: if isinstance(value, bool) or value is None:
return return
@ -147,18 +169,35 @@ def _is_live_status(status: str | None) -> bool:
return status.lower() not in _TERMINAL_INSTANCE_STATUSES return status.lower() not in _TERMINAL_INSTANCE_STATUSES
def _normalized_ref(value: str) -> str:
return "-".join(item for item in re.split(r"[^0-9a-z]+", value.lower()) if item)
def _associated_run_id(instance: Mapping[str, Any], known_run_ids: Sequence[str]) -> str | None: def _associated_run_id(instance: Mapping[str, Any], known_run_ids: Sequence[str]) -> str | None:
for key in ("run_id", "label", "name", "cluster_name"): for key in ("run_id", "label", "name", "cluster_name"):
value = instance.get(key) value = instance.get(key)
if isinstance(value, str) and value: if isinstance(value, str) and value:
if value in known_run_ids: if value in known_run_ids:
return value return value
normalized_value = _normalized_ref(value)
for run_id in known_run_ids: for run_id in known_run_ids:
if run_id and run_id in value: normalized_run_id = _normalized_ref(run_id)
if run_id and (run_id in value or normalized_run_id in normalized_value):
return run_id return run_id
return None return None
def _instance_label(instance: Mapping[str, Any]) -> str | None:
value = _first_present(instance, "label", "name", "cluster_name", "run_id")
return str(value) if value is not None else None
def _matches_ignored_label(label: str | None, prefixes: Sequence[str]) -> bool:
if label is None:
return False
return any(label.startswith(prefix) for prefix in prefixes)
def _first_present(instance: Mapping[str, Any], *keys: str) -> Any: def _first_present(instance: Mapping[str, Any], *keys: str) -> Any:
for key in keys: for key in keys:
if key in instance and instance[key] is not None: if key in instance and instance[key] is not None:

View file

@ -41,6 +41,7 @@ def build_parser() -> argparse.ArgumentParser:
cleanup.add_argument("--vast-instances-json", help="local Vast instances JSON; omit to call Vast API") 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("--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("--destroy-orphans", action="store_true", help="request Vast destruction for live instances missing from Sky")
cleanup.add_argument("--ignore-label-prefix", action="append", default=None, help="do not treat matching external Vast labels as sweep orphans")
cleanup.add_argument("--out") cleanup.add_argument("--out")
cleanup.set_defaults(command="cleanup-reconcile") cleanup.set_defaults(command="cleanup-reconcile")
@ -136,7 +137,7 @@ def main(argv: list[str] | None = None) -> int:
if not api_key: if not api_key:
raise RuntimeError("VAST_API_KEY is required to destroy Vast orphan instances") 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) 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) report = reconcile_cleanup(sky_state=sky_state, vast_instances=instances, ignored_label_prefixes=args.ignore_label_prefix or (), destroy_orphans=args.destroy_orphans, destroy_instance=destroy)
_emit_json(report, args.out) _emit_json(report, args.out)
return 0 return 0
if args.command == "select": if args.command == "select":
@ -676,7 +677,7 @@ def _load_json_file(path: Path) -> Any:
def _load_sky_status() -> Any: def _load_sky_status() -> Any:
process = subprocess.run( process = subprocess.run(
["sky", "status", "--format", "json"], ["sky", "status", "-o", "json"],
check=True, check=True,
capture_output=True, capture_output=True,
text=True, text=True,

View file

@ -46,6 +46,7 @@ class DataConfig:
streaming_queue_max_cases: int = 2 streaming_queue_max_cases: int = 2
streaming_upload_processed: bool = False streaming_upload_processed: bool = False
streaming_upload_batch_size: int = 8 streaming_upload_batch_size: int = 8
streaming_normalization_cases: int | None = None
@dataclass(frozen=True) @dataclass(frozen=True)
@ -236,6 +237,7 @@ def load_training_config(path: str | Path) -> TrainingConfig:
streaming_queue_max_cases=_integer(data_raw, "streaming_queue_max_cases", minimum=1, default=2), 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_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), streaming_upload_batch_size=_integer(data_raw, "streaming_upload_batch_size", minimum=1, default=8),
streaming_normalization_cases=_optional_integer(data_raw, "streaming_normalization_cases", minimum=1),
) )
model = ModelConfig( model = ModelConfig(
type=_choice(_string(model_raw, "type"), _MODEL_TYPES, "model.type"), type=_choice(_string(model_raw, "type"), _MODEL_TYPES, "model.type"),

View file

@ -776,6 +776,10 @@ def _train_public_zip_streaming(
streaming.prepare() streaming.prepare()
stats = streaming.load_or_compute_normalization() stats = streaming.load_or_compute_normalization()
bundle = streaming.schema_bundle() bundle = streaming.schema_bundle()
fast_streaming_start = (
config.data.streaming_normalization_cases is not None
and config.data.streaming_normalization_cases < len(bundle.split.train_ids)
)
writer.write_split_manifest(bundle.split.to_dict()) writer.write_split_manifest(bundle.split.to_dict())
writer.write_json("data_manifest.json", streaming.data_manifest()) writer.write_json("data_manifest.json", streaming.data_manifest())
writer.write_normalization(stats.to_dict()) writer.write_normalization(stats.to_dict())
@ -864,6 +868,22 @@ def _train_public_zip_streaming(
) )
) )
if fast_streaming_start:
record_streaming_metrics(
_log_metrics(
event="streaming_ready" if start_step == 0 else "resume_ready",
step=start_step,
train_loss=None,
val_loss=best_val_loss,
elapsed_seconds=0.0,
lr=_learning_rate(optimizer),
grad_norm=None,
points_per_sec=None,
device=device,
latest_checkpoint=LATEST_CHECKPOINT,
)
)
else:
initial_train = _evaluate_streaming_split( initial_train = _evaluate_streaming_split(
model, model,
streaming, streaming,
@ -983,6 +1003,7 @@ def _train_public_zip_streaming(
latest_checkpoint=LATEST_CHECKPOINT, latest_checkpoint=LATEST_CHECKPOINT,
) )
raise RuntimeError("nonfinite loss") raise RuntimeError("nonfinite loss")
latest_loss_value = float(loss.detach().cpu().item())
loss.backward() loss.backward()
try: try:
grad_norm_tensor = torch.nn.utils.clip_grad_norm_( grad_norm_tensor = torch.nn.utils.clip_grad_norm_(
@ -1038,6 +1059,11 @@ def _train_public_zip_streaming(
) )
if step % log_interval == 0 or step == config.optim.steps: if step % log_interval == 0 or step == config.optim.steps:
if fast_streaming_start:
train_eval = {"loss": latest_loss_value}
val_eval = None
current_metric = latest_loss_value
else:
train_eval = _evaluate_streaming_split( train_eval = _evaluate_streaming_split(
model, model,
streaming, streaming,

View file

@ -562,11 +562,20 @@ class StreamingTrainingData:
return stats return stats
assert self.feature_names is not None assert self.feature_names is not None
assert self.target_names is not None assert self.target_names is not None
self.recorder.emit("normalization_start", phase="normalization", train_cases=len(self.split.train_ids)) normalization_case_limit = self.config.data.streaming_normalization_cases
normalization_case_ids = self.split.train_ids
if normalization_case_limit is not None:
normalization_case_ids = normalization_case_ids[: min(normalization_case_limit, len(normalization_case_ids))]
self.recorder.emit(
"normalization_start",
phase="normalization",
train_cases=len(self.split.train_ids),
normalization_cases=len(normalization_case_ids),
)
started = time.perf_counter() started = time.perf_counter()
accumulator = _StatsAccumulator(feature_names=self.feature_names, target_names=self.target_names) accumulator = _StatsAccumulator(feature_names=self.feature_names, target_names=self.target_names)
rng = np.random.default_rng(self.config.run.seed + _SAMPLE_SEEDS["train"]) rng = np.random.default_rng(self.config.run.seed + _SAMPLE_SEEDS["train"])
for case_id in self.split.train_ids: for case_id in normalization_case_ids:
sample = self.cache.ensure_case(case_id) sample = self.cache.ensure_case(case_id)
self._validate_schema(sample) self._validate_schema(sample)
indices = self._sampling_spec_for_case("train", sample, rng) indices = self._sampling_spec_for_case("train", sample, rng)
@ -584,6 +593,7 @@ class StreamingTrainingData:
"normalization_end", "normalization_end",
phase="normalization", phase="normalization",
train_cases=len(self.split.train_ids), train_cases=len(self.split.train_ids),
normalization_cases=len(normalization_case_ids),
sample_count=accumulator.count, sample_count=accumulator.count,
normalization_runtime_seconds=runtime, normalization_runtime_seconds=runtime,
) )
@ -613,16 +623,30 @@ class StreamingTrainingData:
return self.recorder.to_dict() return self.recorder.to_dict()
def sample_train_batch(self, rng: np.random.Generator, *, batch_size: int, step: int) -> tuple[FloatArray, FloatArray]: def sample_train_batch(self, rng: np.random.Generator, *, batch_size: int, step: int) -> tuple[FloatArray, FloatArray]:
if self.stats is None or self.split is None or self._train_offsets is None: if self.stats is None or self.split is None:
raise RuntimeError("Streaming normalization must be computed before sampling") raise RuntimeError("Streaming normalization must be computed before sampling")
started = time.perf_counter() started = time.perf_counter()
features = np.empty((batch_size, len(self.stats.feature_names)), dtype=np.float32)
targets = np.empty((batch_size, len(self.stats.target_names)), dtype=np.float32)
if self._train_offsets is None:
case_id = self.split.train_ids[int(rng.integers(0, len(self.split.train_ids)))]
sample = self.cache.ensure_case(case_id)
self._validate_schema(sample)
spec = self._sampling_spec_for_case("train", sample, rng)
local_indices = rng.integers(0, spec.count, size=batch_size, dtype=np.int64)
source_indices = _source_indices_for_local(spec, local_indices)
selected_features = sample.features[source_indices]
selected_targets = sample.targets[source_indices]
features[:] = ((selected_features - self.stats.feature_mean) / self.stats.feature_std).astype(np.float32, copy=False)
targets[:] = ((selected_targets - self.stats.target_mean) / self.stats.target_std).astype(np.float32, copy=False)
self._upload_queue.enqueue(sample.source_path)
self.cache.release_case(case_id)
else:
total = int(self._train_offsets[-1]) if self._train_offsets.size else 0 total = int(self._train_offsets[-1]) if self._train_offsets.size else 0
if total <= 0: if total <= 0:
raise ValueError("Streaming train split has no sampled points") raise ValueError("Streaming train split has no sampled points")
global_indices = rng.integers(0, total, size=batch_size) global_indices = rng.integers(0, total, size=batch_size)
case_positions = np.searchsorted(self._train_offsets[1:], global_indices, side="right") case_positions = np.searchsorted(self._train_offsets[1:], global_indices, side="right")
features = np.empty((batch_size, len(self.stats.feature_names)), dtype=np.float32)
targets = np.empty((batch_size, len(self.stats.target_names)), dtype=np.float32)
for case_position in np.unique(case_positions): for case_position in np.unique(case_positions):
mask = case_positions == case_position mask = case_positions == case_position
case_id = self.split.train_ids[int(case_position)] case_id = self.split.train_ids[int(case_position)]
@ -721,6 +745,11 @@ class StreamingTrainingData:
def _build_train_offsets(self) -> None: def _build_train_offsets(self) -> None:
if self.split is None: if self.split is None:
raise RuntimeError("split is missing") raise RuntimeError("split is missing")
normalization_case_limit = self.config.data.streaming_normalization_cases
if normalization_case_limit is not None and normalization_case_limit < len(self.split.train_ids):
self._train_counts = []
self._train_offsets = None
return
counts: list[int] = [] counts: list[int] = []
for case_id in self.split.train_ids: for case_id in self.split.train_ids:
spec = self._sampling_specs["train"].get(case_id) spec = self._sampling_specs["train"].get(case_id)

View file

@ -1,6 +1,9 @@
from __future__ import annotations from __future__ import annotations
import hashlib
import json
import gzip import gzip
import http.client
import sys import sys
import tempfile import tempfile
import types import types
@ -14,7 +17,16 @@ from airfrans_frontier.runtime import remove_pythonpath_entries
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 from airfrans_frontier.raw.bounded_public import prepare_public_airfrans_processed_hf_bounded
from airfrans_frontier.raw.public import (
HttpRangeReader,
_extract_remote_case_members,
_read_zip_central_directory,
_remote_archive_case_members,
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: def write_minimal_airfrans_archive(archive: Path, case_names: list[str]) -> None:
@ -120,6 +132,93 @@ class PublicAirfransDataTests(unittest.TestCase):
self.assertFalse((tmp_path / "streaming_raw" / case_name).exists()) self.assertFalse((tmp_path / "streaming_raw" / case_name).exists())
self.assertGreater(streamed.ranged_bytes_read, 0) self.assertGreater(streamed.ranged_bytes_read, 0)
def test_range_streaming_extracts_contiguous_case_with_one_payload_read(self) -> None:
class CountingRangeReader:
def __init__(self, path: Path) -> None:
self._path = path
self.size = path.stat().st_size
self.bytes_read = 0
self.calls: list[tuple[int, int]] = []
def read_range(self, start: int, length: int) -> bytes:
self.calls.append((start, length))
with self._path.open("rb") as handle:
handle.seek(start)
data = handle.read(length)
self.bytes_read += len(data)
return data
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])
reader = CountingRangeReader(archive)
members = _read_zip_central_directory(reader)
case_members = _remote_archive_case_members(members)[case_name]
reader.calls.clear()
_extract_remote_case_members(reader, case_members, tmp_path / "streaming_raw")
self.assertEqual(len(reader.calls), 1)
self.assertTrue((tmp_path / "streaming_raw" / case_name / "constant" / "transportProperties").is_file())
def test_http_range_reader_retries_timeout_before_failing_run(self) -> None:
class FakeResponse:
status = 206
def __enter__(self):
return self
def __exit__(self, exc_type, exc, traceback):
return False
def read(self) -> bytes:
return b"ok"
reader = HttpRangeReader.__new__(HttpRangeReader)
reader.url = "https://example.test/OF_dataset.zip"
reader.size = 10
reader.bytes_read = 0
with patch("airfrans_frontier.raw.public.urllib.request.urlopen", side_effect=[TimeoutError("timed out"), FakeResponse()]), patch(
"airfrans_frontier.raw.public.time.sleep"
) as sleep:
data = reader.read_range(2, 2)
self.assertEqual(data, b"ok")
self.assertEqual(reader.bytes_read, 2)
sleep.assert_called_once_with(2.0)
def test_http_range_reader_retries_incomplete_body_before_failing_run(self) -> None:
class FakeResponse:
status = 206
def __enter__(self):
return self
def __exit__(self, exc_type, exc, traceback):
return False
def read(self) -> bytes:
return b"ok"
reader = HttpRangeReader.__new__(HttpRangeReader)
reader.url = "https://example.test/OF_dataset.zip"
reader.size = 10
reader.bytes_read = 0
with patch(
"airfrans_frontier.raw.public.urllib.request.urlopen",
side_effect=[http.client.IncompleteRead(b"pa", 2), FakeResponse()],
), patch("airfrans_frontier.raw.public.time.sleep") as sleep:
data = reader.read_range(2, 2)
self.assertEqual(data, b"ok")
self.assertEqual(reader.bytes_read, 2)
sleep.assert_called_once_with(2.0)
def test_prepare_public_hf_streams_archive_before_publish(self) -> None: def test_prepare_public_hf_streams_archive_before_publish(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp) tmp_path = Path(tmp)
@ -157,5 +256,135 @@ class PublicAirfransDataTests(unittest.TestCase):
self.assertEqual(report["processed_case_count"], 1) self.assertEqual(report["processed_case_count"], 1)
def test_bounded_public_hf_uploads_verified_chunks_and_cleans_staging(self) -> None:
uploaded: dict[str, tuple[int, str, bytes]] = {}
commits: list[tuple[str, tuple[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:
self.repo_id = repo_id
def create_commit(self, *, repo_id: str, repo_type: str, operations, commit_message: str):
paths: list[str] = []
for operation in operations:
payload = Path(operation.path_or_fileobj).read_bytes()
uploaded[operation.path_in_repo] = (
len(payload),
hashlib.sha256(payload).hexdigest(),
payload,
)
paths.append(operation.path_in_repo)
commits.append((commit_message, tuple(paths)))
return types.SimpleNamespace(commit_url=f"https://huggingface.co/datasets/{repo_id}/commit/{len(commits)}", oid=str(len(commits)))
def repo_info(self, *, repo_id: str, repo_type: str, files_metadata: bool):
siblings = [
types.SimpleNamespace(rfilename=path, size=size, lfs={"sha256": sha})
for path, (size, sha, _payload) in uploaded.items()
]
return types.SimpleNamespace(siblings=siblings)
def list_repo_files(self, *, repo_id: str, repo_type: str):
return sorted(uploaded)
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 = ["airFoil2D_SST_10.0_5.0_0012", "airFoil2D_SST_11.0_5.0_0012"]
write_minimal_airfrans_archive(archive, case_names)
report = prepare_public_airfrans_processed_hf_bounded(
repo_id="owner/repo",
path_in_repo="processed/full",
work_dir=tmp_path / "work",
output_dir=tmp_path / "staging",
source_url=str(archive),
min_cases=2,
chunk_max_bytes=1,
train_cases=1,
val_cases=1,
test_cases=0,
split_seed=123,
)
state = json.loads((tmp_path / "work" / "bounded_prepare_state.json").read_text())
final_manifest = json.loads((tmp_path / "work" / "chunk_manifests" / "hf_dataset_manifest.json").read_text())
self.assertTrue(report["ok"])
self.assertEqual(report["processed_case_count"], 2)
self.assertEqual(report["chunks_uploaded"], 2)
self.assertEqual(state["phase"], "published")
self.assertEqual(final_manifest["case_count"], 2)
self.assertEqual(final_manifest["feature_names"][0], "x")
self.assertEqual(final_manifest["target_names"], ["velocity_x", "velocity_y", "pressure", "turbulent_viscosity"])
self.assertEqual(final_manifest["split_compatibility"]["train_cases"], 1)
self.assertFalse(any((tmp_path / "staging").glob("*.npz")))
self.assertFalse((tmp_path / "work" / "bounded_raw_scratch").exists())
self.assertIn("processed/full/hf_dataset_manifest.json", uploaded)
self.assertEqual(sum(1 for path in uploaded if path.endswith(".npz")), 2)
self.assertEqual(sum(1 for message, _paths in commits if message.startswith("Upload bounded AirfRANS processed chunk")), 2)
def test_bounded_public_hf_rejects_remote_checksum_mismatch(self) -> None:
uploaded: dict[str, tuple[int, 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:
pass
def create_commit(self, *, repo_id: str, repo_type: str, operations, commit_message: str):
for operation in operations:
payload = Path(operation.path_or_fileobj).read_bytes()
uploaded[operation.path_in_repo] = (len(payload), hashlib.sha256(payload).hexdigest())
return types.SimpleNamespace(commit_url="https://huggingface.co/datasets/owner/repo/commit/bad", oid="bad")
def repo_info(self, *, repo_id: str, repo_type: str, files_metadata: bool):
siblings = [
types.SimpleNamespace(rfilename=path, size=size, lfs={"sha256": "0" * 64})
for path, (size, _sha) in uploaded.items()
]
return types.SimpleNamespace(siblings=siblings)
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"
write_minimal_airfrans_archive(archive, ["airFoil2D_SST_10.0_5.0_0012"])
with self.assertRaisesRegex(RuntimeError, "HF upload verification failed"):
prepare_public_airfrans_processed_hf_bounded(
repo_id="owner/repo",
path_in_repo="processed/full",
work_dir=tmp_path / "work",
output_dir=tmp_path / "staging",
source_url=str(archive),
min_cases=1,
chunk_max_bytes=1024,
train_cases=1,
val_cases=0,
test_cases=0,
split_seed=123,
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()

View file

@ -6,7 +6,7 @@ import json
import tempfile import tempfile
import shutil import shutil
import unittest import unittest
from unittest.mock import patch from unittest.mock import Mock, patch
from pathlib import Path from pathlib import Path
from airfrans_frontier.runtime import remove_pythonpath_entries from airfrans_frontier.runtime import remove_pythonpath_entries
@ -16,6 +16,7 @@ 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.cleanup import reconcile_cleanup
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, _terminal_artifact_names, main as remote_main
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
@ -123,6 +124,104 @@ class VastInstanceCliTests(unittest.TestCase):
self.assertEqual(payload["instances"][0]["id"], 123) self.assertEqual(payload["instances"][0]["id"], 123)
self.assertEqual(payload["instances"][0]["actual_status"], "running") self.assertEqual(payload["instances"][0]["actual_status"], "running")
def test_cleanup_reconcile_uses_current_sky_json_flag(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
vast_path = Path(tmp) / "vast.json"
vast_path.write_text(json.dumps({"instances": []}) + "\n")
stdout = StringIO()
with patch(
"airfrans_frontier.remote.cli.subprocess.run",
return_value=Mock(stdout=json.dumps({"clusters": []})),
) as run, redirect_stdout(stdout):
code = remote_main(["cleanup-reconcile", "--vast-instances-json", str(vast_path)])
self.assertEqual(code, 0)
run.assert_called_once()
self.assertEqual(run.call_args.args[0], ["sky", "status", "-o", "json"])
def test_cleanup_reconcile_treats_sky_node_names_as_live_instances(self) -> None:
destroy_instance = Mock()
report = reconcile_cleanup(
sky_state=[
{
"name": "full_70gb_model_zoo_20260726_film_fourier_inr-20260726T094528Z",
"node_names": "45895903",
"status": "UP",
}
],
vast_instances=[
{
"id": 45895903,
"label": "full-70gb-model-zoo-20260726-film-fourier-inr-20260726t094528z-166d0dab-head",
"actual_status": "running",
"gpu_name": "RTX 4090",
}
],
destroy_orphans=True,
destroy_instance=destroy_instance,
)
self.assertEqual(report["unexpected_live_count"], 0)
self.assertEqual(report["instances"][0]["cleanup_result"], "not_needed")
destroy_instance.assert_not_called()
def test_cleanup_reconcile_ignores_declared_external_vast_labels(self) -> None:
destroy_instance = Mock()
report = reconcile_cleanup(
sky_state=[],
vast_instances=[
{
"id": 45900657,
"label": "mvp-1785062905827-4",
"actual_status": "running",
"gpu_name": "Titan RTX",
},
{
"id": 45901000,
"label": "full-70gb-model-zoo-20260726-film-fourier-inr-canary-head",
"actual_status": "running",
"gpu_name": "RTX 4090",
},
],
ignored_label_prefixes=("mvp-",),
destroy_orphans=True,
destroy_instance=destroy_instance,
)
self.assertEqual(report["ignored_external_count"], 1)
self.assertEqual(report["unexpected_live_count"], 1)
self.assertEqual(report["instances"][0]["cleanup_result"], "ignored")
self.assertEqual(report["instances"][0]["label"], "mvp-1785062905827-4")
destroy_instance.assert_called_once_with(45901000)
def test_cleanup_reconcile_matches_hyphenated_vast_labels_to_sky_names(self) -> None:
destroy_instance = Mock()
report = reconcile_cleanup(
sky_state=[
{
"name": "full_70gb_model_zoo_20260726_raster_fno_unet-20260726T095257Z",
"status": "INIT",
}
],
vast_instances=[
{
"id": 45896431,
"label": "full-70gb-model-zoo-20260726-raster-fno-unet-20260726t095257z-166d0dab-head",
"actual_status": "loading",
"gpu_name": "RTX 4090",
}
],
destroy_orphans=True,
destroy_instance=destroy_instance,
)
self.assertEqual(report["unexpected_live_count"], 0)
self.assertEqual(
report["instances"][0]["associated_run_id"],
"full_70gb_model_zoo_20260726_raster_fno_unet-20260726T095257Z",
)
destroy_instance.assert_not_called()
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:

View file

@ -67,6 +67,7 @@ def write_streaming_config(
low_water_bytes: int = 16 * 1024 * 1024, low_water_bytes: int = 16 * 1024 * 1024,
upload_processed: bool = False, upload_processed: bool = False,
upload_batch_size: int = 1, upload_batch_size: int = 1,
normalization_cases: int | None = None,
) -> None: ) -> None:
path.write_text( path.write_text(
f""" f"""
@ -92,6 +93,7 @@ streaming_cache_low_water_bytes = {low_water_bytes}
streaming_queue_max_cases = 1 streaming_queue_max_cases = 1
streaming_upload_processed = {str(upload_processed).lower()} streaming_upload_processed = {str(upload_processed).lower()}
streaming_upload_batch_size = {upload_batch_size} streaming_upload_batch_size = {upload_batch_size}
{f"streaming_normalization_cases = {normalization_cases}" if normalization_cases is not None else ""}
hf_repo_id = "owner/airfrans-processed" hf_repo_id = "owner/airfrans-processed"
hf_repo_type = "dataset" hf_repo_type = "dataset"
hf_path_prefix = "processed/full" hf_path_prefix = "processed/full"
@ -180,6 +182,39 @@ class FullDataBackpressureStreamingTests(unittest.TestCase):
self.assertLessEqual(processed_cases, selected_cases) self.assertLessEqual(processed_cases, selected_cases)
self.assertFalse(any((tmp_path / "cache" / "_raw").glob("airFoil2D_*"))) self.assertFalse(any((tmp_path / "cache" / "_raw").glob("airFoil2D_*")))
def test_streaming_fast_start_reaches_first_gpu_batch_before_all_train_cases_are_processed(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(6)]
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,
train_cases=4,
val_cases=1,
test_cases=1,
steps=1,
normalization_cases=1,
)
result = train(load_training_config(config_path))
events = read_events(result.run_dir)
first_gpu = next(index for index, event in enumerate(events) if event["event"] == "first_gpu_batch_consumed")
processed_before_gpu = {
str(event["case_id"])
for event in events[:first_gpu]
if event["event"] == "processing_end"
}
self.assertLess(len(processed_before_gpu), 4)
normalization_end = next(event for event in events if event["event"] == "normalization_end")
self.assertEqual(normalization_end["normalization_cases"], 1)
def test_backpressure_pauses_resumes_and_bounds_cache_with_inflight_slack(self) -> None: def test_backpressure_pauses_resumes_and_bounds_cache_with_inflight_slack(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
tmp_path = Path(tmp) tmp_path = Path(tmp)