Skip to content
Open
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
1 change: 1 addition & 0 deletions .gitattributes
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
**/kernel_workloads/**/*.safetensors filter=lfs diff=lfs merge=lfs -text
2 changes: 1 addition & 1 deletion modeling/transformers/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ build-backend = "setuptools.build_meta"
name = "tilegym-hf-bench"
version = "0.1.0"
description = "Hugging Face inference benchmarks and profiler tooling for TileGym"
requires-python = ">=3.10"
requires-python = ">=3.12"
dependencies = [
"accelerate==1.13.0",
"cuda-bindings>=13.2.0",
Expand Down
1,511 changes: 57 additions & 1,454 deletions modeling/transformers/uv.lock

Large diffs are not rendered by default.

10 changes: 10 additions & 0 deletions src/tilegym/kernel_inventory/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,8 @@
from tilegym.kernel_inventory.composition import DefinitionCompositionError
from tilegym.kernel_inventory.composition import validate_definition_composition
from tilegym.kernel_inventory.layout import iter_inventory_json_paths
from tilegym.kernel_inventory.layout import iter_inventory_workload_paths
from tilegym.kernel_inventory.layout import mirrored_workload_path
from tilegym.kernel_inventory.layout import solution_paths_for_definition
from tilegym.kernel_inventory.return_contract import ReturnContractError
from tilegym.kernel_inventory.return_contract import instrument_reference_returns
Expand All @@ -39,6 +41,14 @@
from tilegym.kernel_inventory.source_contract import SourceContractError
from tilegym.kernel_inventory.source_contract import resolve_repo_relative_path
from tilegym.kernel_inventory.source_contract import validate_reference_source_contract
from tilegym.kernel_inventory.workloads import KernelWorkloadError
from tilegym.kernel_inventory.workloads import WorkloadRecord
from tilegym.kernel_inventory.workloads import iter_workload_records
from tilegym.kernel_inventory.workloads import load_workload_jsonl
from tilegym.kernel_inventory.workloads import materialize_workload_inputs
from tilegym.kernel_inventory.workloads import resolve_torch_dtype
from tilegym.kernel_inventory.workloads import validate_workload_against_definition
from tilegym.kernel_inventory.workloads import validate_workload_catalog

DEFINITION_SCHEMA_URL = (
"https://github.com/flashinfer-ai/flashinfer-bench/blob/main/docs/flashinfer-trace/definition.mdx"
Expand Down
222 changes: 222 additions & 0 deletions src/tilegym/kernel_inventory/_workload_schema_compat.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,222 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# SPDX-License-Identifier: MIT

"""Public compatibility models for the Kernel Factory Workload contract.

Internal builds validate against the pinned canonical schema directly.
Public TileGym releases cannot depend on that private package, so this module
implements the documented Workload subset used by the checked-in inventory.
Project-specific file and Definition compatibility rules intentionally live
outside these schema models.
"""

from __future__ import annotations

import math
from enum import Enum
from typing import Annotated
from typing import Any
from typing import Literal
from typing import TypeAlias

from pydantic import BaseModel
from pydantic import Field
from pydantic import StrictBool
from pydantic import StrictFloat
from pydantic import StrictInt
from pydantic import field_validator
from pydantic import model_serializer
from pydantic import model_validator


class EvalMode(str, Enum):
"""Evaluation phases requested for one Workload."""

FULL = "full"
CORRECTNESS_ONLY = "correctness_only"
BENCHMARK_ONLY = "benchmark_only"


class SamplingStrategy(str, Enum):
"""Sampling strategy for correctness-only dynamic axes."""

RANDOM = "random"
LINEAR = "linear"


class DynamicAxis(BaseModel):
"""Correctness-only dynamic shape sampling policy for one axis."""

min: int | str
max: int | str
multiple_of: int | str = 1
sampling_strategy: SamplingStrategy = SamplingStrategy.RANDOM
intermediate_samples: int = Field(default=0, ge=0)

@field_validator("min", "max")
@classmethod
def _validate_bound(cls, value: int | str) -> int | str:
if isinstance(value, bool) or isinstance(value, int) and value < 0:
raise ValueError("dynamic axis bounds must be non-negative integers or const-axis names")
if isinstance(value, str) and not value:
raise ValueError("dynamic axis const-axis names must be non-empty")
return value

@field_validator("multiple_of")
@classmethod
def _validate_multiple_of(cls, value: int | str) -> int | str:
if isinstance(value, bool) or isinstance(value, int) and value <= 0:
raise ValueError("dynamic axis multiple_of must be positive")
if isinstance(value, str) and not value:
raise ValueError("dynamic axis multiple_of const-axis name must be non-empty")
return value

@model_validator(mode="after")
def _validate_numeric_bounds(self) -> "DynamicAxis":
if isinstance(self.min, int) and isinstance(self.max, int) and self.min > self.max:
raise ValueError("dynamic axis min must not exceed max")
return self


class RandomInput(BaseModel):
"""Random tensor input descriptor."""

type: Literal["random"] = "random"


ScalarValue: TypeAlias = StrictInt | StrictFloat | StrictBool


class ScalarInput(BaseModel):
"""Python numeric scalar input descriptor."""

type: Literal["scalar"] = "scalar"
value: ScalarValue


class SafetensorsShard(BaseModel):
"""One rank-local tensor locator."""

path: str = Field(min_length=1)
tensor_key: str = Field(min_length=1)


class SafetensorsInput(BaseModel):
"""Replicated or rank-sharded safetensors input descriptor."""

type: Literal["safetensors"] = "safetensors"
path: str | None = None
tensor_key: str | None = None
shards: list[SafetensorsShard] | None = None

@model_validator(mode="after")
def _validate_locator(self) -> "SafetensorsInput":
has_path = self.path is not None
has_key = self.tensor_key is not None
has_shards = self.shards is not None
if has_path != has_key:
raise ValueError("safetensors path and tensor_key must be specified together")
if has_shards and (has_path or has_key):
raise ValueError("safetensors replicated locator and shards are mutually exclusive")
if not has_shards and not has_path:
raise ValueError("safetensors input requires path/tensor_key or shards")
if has_path and (not self.path or not self.tensor_key):
raise ValueError("safetensors path and tensor_key must be non-empty")
if has_shards and not self.shards:
raise ValueError("safetensors shards must be non-empty")
return self

@model_serializer(mode="wrap")
def _serialize_without_unused_locators(self, handler: Any) -> dict[str, Any]:
return {key: value for key, value in handler(self).items() if value is not None}


class NullInput(BaseModel):
"""Absent optional input descriptor."""

type: Literal["null"] = "null"


class StringInput(BaseModel):
"""Python string input descriptor."""

type: Literal["string"] = "string"
value: str


class CustomInput(BaseModel):
"""Definition-provided custom input descriptor."""

type: Literal["custom"] = "custom"


InputSpec: TypeAlias = Annotated[
RandomInput | ScalarInput | SafetensorsInput | NullInput | StringInput | CustomInput,
Field(discriminator="type"),
]


class ToleranceSpec(BaseModel):
"""Numerical correctness bounds for one Workload."""

max_atol: float = Field(default=0.01, ge=0.0, allow_inf_nan=False)
max_rtol: float = Field(default=0.01, ge=0.0, allow_inf_nan=False)
required_matched_ratio: float = Field(default=0.99, ge=0.0, le=1.0, allow_inf_nan=False)
max_error_cap: float | None = Field(default=None, ge=0.0, allow_inf_nan=False)
allow_negative_inf: bool = False


def _validate_finite_json(value: Any, path: str = "custom_correctness_kwargs") -> None:
if value is None or isinstance(value, str | bool | int):
return
if isinstance(value, float):
if not math.isfinite(value):
raise ValueError(f"{path} must contain only finite JSON numbers")
return
if isinstance(value, list):
for index, item in enumerate(value):
_validate_finite_json(item, f"{path}[{index}]")
return
if isinstance(value, dict):
for key, item in value.items():
if not isinstance(key, str):
raise ValueError(f"{path} object keys must be strings")
_validate_finite_json(item, f"{path}.{key}")
return
raise ValueError(f"{path} contains a non-JSON value")


class Workload(BaseModel):
"""Concrete Kernel Factory-compatible workload configuration."""

axes: dict[str, Annotated[int, Field(ge=0)] | DynamicAxis]
inputs: dict[str, InputSpec]
uuid: str = Field(min_length=1)
tolerance: ToleranceSpec = Field(default_factory=ToleranceSpec)
custom_correctness_kwargs: dict[str, Any] = Field(default_factory=dict)
eval_mode: EvalMode = EvalMode.FULL
weight: float | None = Field(default=None, gt=0.0, allow_inf_nan=False)

@field_validator("axes")
@classmethod
def _validate_axis_names(cls, axes: dict[str, int | DynamicAxis]) -> dict[str, int | DynamicAxis]:
if any(not name for name in axes):
raise ValueError("workload axis names must be non-empty")
return axes

@field_validator("custom_correctness_kwargs")
@classmethod
def _validate_custom_correctness_kwargs(cls, value: dict[str, Any]) -> dict[str, Any]:
_validate_finite_json(value)
return value

@model_validator(mode="after")
def _validate_cross_field_contract(self) -> "Workload":
has_dynamic_axes = any(isinstance(value, DynamicAxis) for value in self.axes.values())
if has_dynamic_axes and self.eval_mode is not EvalMode.CORRECTNESS_ONLY:
raise ValueError("dynamic axes require eval_mode='correctness_only'")
custom_count = sum(isinstance(value, CustomInput) for value in self.inputs.values())
if custom_count and custom_count != len(self.inputs):
raise ValueError("custom inputs cannot be mixed with other input descriptor types")
return self
Loading