Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 55 additions & 9 deletions include/pops/runtime/config/platform_manifest.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -225,12 +225,12 @@ inline void require_same(const std::string& field, const CapabilityProof& expect
throw ContractError(field, field + " mismatch between artifact and runtime backend");
}

inline const CapabilityProof& capability(const RuntimeBackendManifest& backend,
const std::string& name) {
const auto found = backend.capabilities.find(name);
if (found == backend.capabilities.end())
template <class Manifest>
inline const CapabilityProof& capability(const Manifest& manifest, const std::string& name) {
const auto found = manifest.capabilities.find(name);
if (found == manifest.capabilities.end())
throw ContractError("capabilities." + name,
"runtime backend omitted required capability proof " + name);
"platform/runtime manifest omitted required capability proof " + name);
return found->second;
}

Expand All @@ -248,6 +248,18 @@ inline void validate_descriptor(const FieldViewDescriptor& view) {
if (std::any_of(view.ghosts.begin(), view.ghosts.end(),
[](const auto& pair) { return pair.first < 0 || pair.second < 0; }))
throw ContractError("field.ghosts", "field ghost widths must be non-negative");
for (std::size_t axis = 0; axis < rank; ++axis) {
const auto lower = static_cast<std::size_t>(view.ghosts[axis].first);
const auto upper = static_cast<std::size_t>(view.ghosts[axis].second);
if (lower >= view.extents[axis] || upper >= view.extents[axis] - lower)
throw ContractError("field.ghosts",
"field ghost widths must leave a positive interior extent");
}
if (view.centering.empty() || view.scalar.empty() || view.memory_space.empty() ||
view.patch.empty() || view.layout.empty() || view.ownership.empty())
throw ContractError("field.metadata",
"field centering, scalar, memory space, patch, layout and ownership "
"must be non-empty");
}

template <class Value>
Expand Down Expand Up @@ -281,33 +293,67 @@ inline void validate_launch(const PlatformManifest& platform, const ExecutionCon
!context.device.has_handle)
throw ContractError("device", "non-host execution requires an explicit handle");

for (const std::string name :
{"dimensions", "centerings", "scalars", "layouts", "ownership", "generic_field_view"})
require_same("capabilities." + name, capability(platform, name), capability(backend, name));
const auto& generic_field_view =
require(capability(backend, "generic_field_view"), "runtime.capabilities.generic_field_view");
if (generic_field_view.kind() != CanonicalValue::Kind::kBool || !generic_field_view.boolean())
throw ContractError("generic_field_view",
"runtime does not prove the generic field-view launch contract");
const auto dimensions =
require_int_set(capability(backend, "dimensions"), "runtime.capabilities.dimensions");
const auto centerings =
require_text_set(capability(backend, "centerings"), "runtime.capabilities.centerings");
const auto scalars =
require_text_set(capability(backend, "scalars"), "runtime.capabilities.scalars");
const auto layouts =
require_text_set(capability(backend, "layouts"), "runtime.capabilities.layouts");
const auto ownership =
require_text_set(capability(backend, "ownership"), "runtime.capabilities.ownership");
const auto memories = require_text_set(backend.memory_spaces, "runtime.memory_spaces");
for (const auto& view : fields) {
std::vector<std::string> field_names;
field_names.reserve(fields.size());
std::vector<std::string> expected_names;
expected_names.reserve(expected.size());
const auto validate_unique_name = [](const FieldViewDescriptor& view,
std::vector<std::string>& names, const std::string& owner) {
if (std::find(names.begin(), names.end(), view.name) != names.end())
throw ContractError("field." + view.name,
owner + " field descriptors must have unique names");
names.push_back(view.name);
};
const auto validate_capabilities = [&](const FieldViewDescriptor& view) {
validate_descriptor(view);
require_member("dimension", view.dimension, dimensions);
require_member("centering", view.centering, centerings);
require_member("scalar", view.scalar, scalars);
require_member("memory_space", view.memory_space, memories);
require_member("layout", view.layout, layouts);
require_member("ownership", view.ownership, ownership);
};
for (const auto& view : fields) {
validate_unique_name(view, field_names, "launch");
validate_capabilities(view);
if (view.scalar != context.datatype.identity)
throw ContractError("datatype", "field scalar and ExecutionContext datatype differ");
const auto wanted = std::find_if(expected.begin(), expected.end(),
[&](const auto& item) { return item.name == view.name; });
if (wanted != expected.end() &&
(view.dimension != wanted->dimension || view.extents != wanted->extents ||
view.centering != wanted->centering || view.scalar != wanted->scalar ||
view.memory_space != wanted->memory_space))
view.strides != wanted->strides || view.centering != wanted->centering ||
view.ghosts != wanted->ghosts || view.scalar != wanted->scalar ||
view.memory_space != wanted->memory_space || view.patch != wanted->patch ||
view.layout != wanted->layout || view.ownership != wanted->ownership))
throw ContractError("field." + view.name, "field descriptor does not match launch contract");
}
for (const auto& wanted : expected)
for (const auto& wanted : expected) {
validate_unique_name(wanted, expected_names, "expected");
validate_capabilities(wanted);
if (std::none_of(fields.begin(), fields.end(),
[&](const auto& view) { return view.name == wanted.name; }))
throw ContractError("field." + wanted.name, "required field descriptor is missing");
}
}

template <class Kernel>
Expand Down
132 changes: 119 additions & 13 deletions python/pops/_platform_contracts.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,26 @@
_CENTERINGS = frozenset({"cell", "node", "face_x", "face_y", "face_z"})
_LAYOUTS = frozenset({"right", "left", "strided"})
_OWNERSHIP = frozenset({"borrowed", "owned", "shared"})
_FIELD_CAPABILITIES = (
"dimensions",
"centerings",
"scalars",
"layouts",
"ownership",
"generic_field_view",
)
_EXACT_FIELD_ATTRIBUTES = (
"dimension",
"extents",
"strides",
"centering",
"ghosts",
"scalar",
"memory_space",
"patch",
"layout",
"ownership",
)
_STD_YEARS = {"11": "201103", "14": "201402", "17": "201703",
"20": "202002", "23": "202302"}

Expand Down Expand Up @@ -324,6 +344,10 @@ def __post_init__(self) -> None:
len(pair) != 2 or any(isinstance(item, bool) or not isinstance(item, int)
or item < 0 for item in pair) for pair in ghosts):
raise ValueError("FieldViewDescriptor.ghosts must contain one non-negative pair per axis")
if any(lower >= extent or upper >= extent - lower
for extent, (lower, upper) in zip(self.extents, ghosts, strict=True)):
raise ValueError(
"FieldViewDescriptor.ghosts must leave a positive interior extent on every axis")
object.__setattr__(self, "ghosts", ghosts)
if self.centering not in _CENTERINGS:
raise ValueError("unsupported field centering %r" % self.centering)
Expand Down Expand Up @@ -411,35 +435,71 @@ def _validate_launch_facts(platform: PlatformManifest, context: ExecutionContext
for name in ("storage", "compute", "accumulation", "reduction"):
_require_same("precision.%s" % name, getattr(platform.precision, name),
getattr(backend.precision, name))
supported_dimensions = tuple(backend.capabilities["dimensions"].require(
for name in _FIELD_CAPABILITIES:
_require_same(
"capabilities.%s" % name,
_field_capability(platform, name, owner="artifact"),
_field_capability(backend, name, owner="runtime"),
)
generic_field_view = _field_capability(
backend, "generic_field_view", owner="runtime").require(
"runtime.capabilities.generic_field_view")
if type(generic_field_view) is not bool or not generic_field_view:
raise PlatformContractError(
"runtime does not prove the generic field-view launch contract",
field="generic_field_view", expected=True, actual=generic_field_view)
supported_dimensions = tuple(_field_capability(
backend, "dimensions", owner="runtime").require(
"runtime.capabilities.dimensions"))
supported_centerings = tuple(backend.capabilities["centerings"].require(
supported_centerings = tuple(_field_capability(
backend, "centerings", owner="runtime").require(
"runtime.capabilities.centerings"))
supported_scalars = tuple(backend.capabilities["scalars"].require(
supported_scalars = tuple(_field_capability(
backend, "scalars", owner="runtime").require(
"runtime.capabilities.scalars"))
supported_layouts = tuple(_field_capability(
backend, "layouts", owner="runtime").require(
"runtime.capabilities.layouts"))
supported_ownership = tuple(_field_capability(
backend, "ownership", owner="runtime").require(
"runtime.capabilities.ownership"))
supported_memory = tuple(backend.memory_spaces.require("runtime.memory_spaces"))
actual = tuple(fields)
expected = {item.name: item for item in expected_fields}
if len(expected) != len(tuple(expected_fields)):
raise ValueError("expected field names must be unique")
required = tuple(expected_fields)
_require_unique_field_names(actual, owner="launch")
_require_unique_field_names(required, owner="expected")
expected = {item.name: item for item in required}
for view in actual:
if type(view) is not FieldViewDescriptor:
raise TypeError("fields must contain exact FieldViewDescriptor values")
_require_field_capability(view, "dimension", view.dimension, supported_dimensions)
_require_field_capability(view, "centering", view.centering, supported_centerings)
_require_field_capability(view, "scalar", view.scalar, supported_scalars)
_require_field_capability(view, "memory_space", view.memory_space, supported_memory)
_validate_field_capabilities(
view,
dimensions=supported_dimensions,
centerings=supported_centerings,
scalars=supported_scalars,
memory_spaces=supported_memory,
layouts=supported_layouts,
ownership=supported_ownership,
)
if view.scalar != context.datatype.identity:
raise PlatformContractError(
"field scalar does not match ExecutionContext datatype", field="datatype",
expected=view.scalar, actual=context.datatype.identity)
requirement = expected.get(view.name)
if requirement is not None:
for name in ("dimension", "extents", "centering", "scalar", "memory_space"):
for name in _EXACT_FIELD_ATTRIBUTES:
if getattr(view, name) != getattr(requirement, name):
raise PlatformContractError(
"field %r %s mismatch" % (view.name, name), field=name,
expected=getattr(requirement, name), actual=getattr(view, name))
for view in required:
_validate_field_capabilities(
view,
dimensions=supported_dimensions,
centerings=supported_centerings,
scalars=supported_scalars,
memory_spaces=supported_memory,
layouts=supported_layouts,
ownership=supported_ownership,
)
missing = sorted(set(expected) - {item.name for item in actual})
if missing:
raise PlatformContractError("required field view(s) are missing: %s" % missing,
Expand Down Expand Up @@ -502,6 +562,12 @@ def validate_component_runtime(platform: PlatformManifest,
_require_same(
"capabilities.%s" % name,
platform.capabilities[name], runtime.capabilities[name])
for name in _FIELD_CAPABILITIES:
_require_same(
"capabilities.%s" % name,
_field_capability(platform, name, owner="component"),
_field_capability(runtime, name, owner="runtime"),
)
expected_abi = platform.abi.require("component.abi")
actual_abi = runtime.abi.require("runtime.abi")
if expected_abi != actual_abi:
Expand All @@ -522,6 +588,46 @@ def _require_field_capability(view: FieldViewDescriptor, field_name: str,
expected=supported, actual=value)


def _field_capability(manifest: PlatformManifest | RuntimeBackendManifest, name: str,
*, owner: str) -> CapabilityProof:
proof = manifest.capabilities.get(name)
if proof is None:
raise PlatformContractError(
"%s omitted required field-view capability %r" % (owner, name),
field="capabilities.%s" % name, expected="explicit proof", actual=None)
return proof


def _require_unique_field_names(fields: tuple[FieldViewDescriptor, ...], *, owner: str) -> None:
names: set[str] = set()
for view in fields:
if type(view) is not FieldViewDescriptor:
raise TypeError("%s fields must contain exact FieldViewDescriptor values" % owner)
if view.name in names:
raise PlatformContractError(
"%s field descriptors contain duplicate name %r" % (owner, view.name),
field="fields.%s" % view.name, expected="unique name", actual=view.name)
names.add(view.name)


def _validate_field_capabilities(
view: FieldViewDescriptor,
*,
dimensions: tuple[Any, ...],
centerings: tuple[Any, ...],
scalars: tuple[Any, ...],
memory_spaces: tuple[Any, ...],
layouts: tuple[Any, ...],
ownership: tuple[Any, ...],
) -> None:
_require_field_capability(view, "dimension", view.dimension, dimensions)
_require_field_capability(view, "centering", view.centering, centerings)
_require_field_capability(view, "scalar", view.scalar, scalars)
_require_field_capability(view, "memory_space", view.memory_space, memory_spaces)
_require_field_capability(view, "layout", view.layout, layouts)
_require_field_capability(view, "ownership", view.ownership, ownership)


def launch_checked(platform: PlatformManifest, context: ExecutionContext,
fields: Sequence[FieldViewDescriptor], kernel: Callable[..., Any],
*, expected_fields: Sequence[FieldViewDescriptor] = ()) -> Any:
Expand Down
65 changes: 61 additions & 4 deletions tests/cpp/unit/runtime/test_platform_manifest.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -72,9 +72,8 @@ TEST(PlatformManifest, FieldAndCommunicatorMismatchesRefuseBeforeKernel) {
int launches = 0;
auto kernel = [&](const auto&, const auto&) { return ++launches; };
const auto required = field();
for (int variant = 0; variant < 5; ++variant) {
for (int variant = 0; variant < 9; ++variant) {
auto actual = field();
auto execution = context();
if (variant == 0)
actual.centering = "node";
else if (variant == 1)
Expand All @@ -83,13 +82,71 @@ TEST(PlatformManifest, FieldAndCommunicatorMismatchesRefuseBeforeKernel) {
actual.extents = {15, 12};
else if (variant == 3)
actual.memory_space = "device";
else if (variant == 4)
actual.strides = {1, 16};
else if (variant == 5)
actual.ghosts = {{1, 0}, {0, 0}};
else if (variant == 6)
actual.patch = "patch-1";
else if (variant == 7)
actual.layout = "left";
else
execution.communicator.identity = "comm:wrong";
actual.ownership = "owned";
EXPECT_THROW(
pops::platform::launch_checked(platform(), execution, {actual}, kernel, {required}),
pops::platform::launch_checked(platform(), context(), {actual}, kernel, {required}),
pops::platform::ContractError);
EXPECT_EQ(launches, 0);
}
auto execution = context();
execution.communicator.identity = "comm:wrong";
EXPECT_THROW(pops::platform::launch_checked(platform(), execution, {field()}, kernel, {required}),
pops::platform::ContractError);
EXPECT_EQ(launches, 0);
}

TEST(PlatformManifest, FieldCapabilitiesAndNamesFailClosed) {
int launches = 0;
auto kernel = [&](const auto&, const auto&) { return ++launches; };

auto missing = platform();
missing.capabilities.erase("ownership");
EXPECT_THROW(pops::platform::launch_checked(missing, context(), {field()}, kernel),
pops::platform::ContractError);

auto unsupported = platform();
unsupported.capabilities["layouts"] = pops::platform::prove_text_set({"left"}, "test");
auto unsupported_context = context();
unsupported_context.backend.capabilities["layouts"] =
pops::platform::prove_text_set({"left"}, "test");
EXPECT_THROW(pops::platform::launch_checked(unsupported, unsupported_context, {field()}, kernel),
pops::platform::ContractError);

auto disabled = platform();
disabled.capabilities["generic_field_view"] = pops::platform::prove_bool(false, "test");
auto disabled_context = context();
disabled_context.backend.capabilities["generic_field_view"] =
pops::platform::prove_bool(false, "test");
EXPECT_THROW(pops::platform::launch_checked(disabled, disabled_context, {field()}, kernel),
pops::platform::ContractError);

EXPECT_THROW(pops::platform::launch_checked(platform(), context(), {field(), field()}, kernel),
pops::platform::ContractError);
EXPECT_THROW(
pops::platform::launch_checked(platform(), context(), {field()}, kernel, {field(), field()}),
pops::platform::ContractError);
EXPECT_EQ(launches, 0);
}

TEST(PlatformManifest, FieldGhostsMustLeavePositiveInterior) {
auto hidden = field();
hidden.ghosts = {{16, 0}, {0, 0}};
EXPECT_THROW(pops::platform::validate_launch(platform(), context(), {hidden}),
pops::platform::ContractError);

hidden = field();
hidden.ghosts = {{8, 8}, {0, 0}};
EXPECT_THROW(pops::platform::validate_launch(platform(), context(), {hidden}),
pops::platform::ContractError);
}

TEST(PlatformManifest, GenericTwoDimensionalDoubleRouteLaunches) {
Expand Down
Loading
Loading