poincare-hyper / tests /test_well_adapter.py
DHDRL's picture
Upload 30 files
ae73c7f verified
Raw
History Blame Contribute Delete
3.65 kB
"""
Tests for well_adapter.py, built against the_well package
interface (verified via inspect.getsource on the installed
library): WellDataset samples are dicts with input_fields/
output_fields shaped (T, H, W, C) -- channels last.
"""
from __future__ import annotations
import os
import sys
import pytest
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from src.well_adapter import well_sample_to_fields, WellStreamAdapter
from src.provenance import SchemaValidationError
def make_well_sample(Ti=4, To=4, H=32, W=32, C=2):
return {
"input_fields": torch.randn(Ti, H, W, C),
"output_fields": torch.randn(To, H, W, C),
"constant_fields": torch.randn(H, W, 1), # present in real samples, must be ignored
}
def test_well_sample_to_fields_converts_layout_correctly():
sample = make_well_sample(Ti=4, To=4, H=32, W=32, C=2)
fields = well_sample_to_fields(sample, include_output=True)
assert fields.shape == (8, 2, 32, 32) # (Ti+To, C, H, W)
def test_well_sample_to_fields_input_only():
sample = make_well_sample(Ti=4, To=4, H=16, W=16, C=11)
fields = well_sample_to_fields(sample, include_output=False)
assert fields.shape == (4, 11, 16, 16)
def test_well_sample_to_fields_actually_permutes_not_just_reshapes():
Ti, H, W, C = 1, 3, 4, 2
sample = {
"input_fields": torch.arange(Ti * H * W * C).float().reshape(Ti, H, W, C),
"output_fields": torch.zeros(0, H, W, C),
}
fields = well_sample_to_fields(sample, include_output=True)
assert fields.shape == (1, 2, 3, 4)
# channel 0 at (h=0,w=0) in source should land at fields[0,0,0,0]
original = sample["input_fields"]
assert torch.equal(fields[0, 0], original[0, :, :, 0])
assert torch.equal(fields[0, 1], original[0, :, :, 1])
def test_well_sample_to_fields_raises_on_missing_input_fields():
with pytest.raises(SchemaValidationError) as exc_info:
well_sample_to_fields({"output_fields": torch.randn(2, 8, 8, 2)})
assert exc_info.value.outcome_code == "WELL_SAMPLE_MISSING_INPUT_FIELDS"
def test_well_sample_to_fields_raises_on_missing_output_fields_when_requested():
with pytest.raises(SchemaValidationError) as exc_info:
well_sample_to_fields({"input_fields": torch.randn(2, 8, 8, 2)}, include_output=True)
assert exc_info.value.outcome_code == "WELL_SAMPLE_MISSING_OUTPUT_FIELDS"
def test_well_sample_to_fields_raises_on_wrong_ndim():
with pytest.raises(SchemaValidationError) as exc_info:
well_sample_to_fields({"input_fields": torch.randn(2, 8, 8)}, include_output=False) # 3D, missing C
assert exc_info.value.outcome_code == "WELL_SAMPLE_UNEXPECTED_SHAPE"
def test_well_sample_to_fields_raises_on_mismatched_hwc():
sample = {
"input_fields": torch.randn(4, 32, 32, 2),
"output_fields": torch.randn(4, 16, 16, 2), # different H,W
}
with pytest.raises(SchemaValidationError) as exc_info:
well_sample_to_fields(sample, include_output=True)
assert exc_info.value.outcome_code == "WELL_SAMPLE_SHAPE_MISMATCH"
def test_well_stream_adapter_wraps_correctly():
class FakeWellStream:
provenance = "REAL_STREAMED"
def __len__(self): return 3
def __getitem__(self, idx): return make_well_sample(C=5)
adapter = WellStreamAdapter(FakeWellStream())
assert len(adapter) == 3
assert adapter.provenance == "REAL_STREAMED"
item = adapter[0]
assert set(item.keys()) >= {"fields", "idx"}
assert item["fields"].shape == (8, 5, 32, 32)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))