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
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
{
"export": {
"opset_version": 17,
"batch_size": 1,
"export_params": true,
"do_constant_folding": true,
"verbose": false,
"dynamo": false,
"enable_hierarchy_tags": true,
"clean_onnx": false,
"hierarchy_tag_format": "full",
"input_tensors": [
{
"name": "input_ids",
"dtype": "int32",
"shape": [
1,
512
],
"value_range": [
0,
128100
]
},
{
"name": "attention_mask",
"dtype": "int32",
"shape": [
1,
512
],
"value_range": [
0,
2
]
}
],
"output_tensors": [
{
"name": "logits"
}
],
"compatibility": {
"transformers_attention": "eager"
}
},
"optim": {},
"quant": {
"mode": "fp16",
"samples": 10,
"calibration_method": "minmax",
"weight_type": "uint8",
"activation_type": "uint8",
"per_channel": false,
"symmetric": false,
"weight_symmetric": null,
"activation_symmetric": null,
"save_calibration": false,
"distribution": "uniform",
"seed": null,
"calibration_load_path": null,
"calibration_save_path": null,
"op_types_to_quantize": null,
"nodes_to_exclude": null,
"task": "reranking",
"model_id": "mixedbread-ai/mxbai-rerank-base-v1",
"model_type": "deberta-v2",
"fp16_keep_io_types": true,
"fp16_op_block_list": null,
"fp16_nodes_to_exclude": [
"InsertedPrecisionFreeCast_/deberta/embeddings/LayerNorm/LayerNormalization_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.0/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.0/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.0/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.0/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.0/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.0/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.1/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.1/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.1/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.1/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.1/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.1/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.2/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.2/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.2/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.2/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.2/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.2/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.3/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.3/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.3/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.3/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.3/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.3/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.4/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.4/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.4/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.4/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.4/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.4/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.5/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.5/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.5/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.5/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.5/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.5/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.6/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.6/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.6/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.6/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.6/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.6/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.7/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.7/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.7/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.7/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.7/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.7/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.8/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.8/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.8/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.8/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.8/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.8/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.9/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.9/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.9/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.9/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.9/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.9/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.10/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.10/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.10/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.10/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.10/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.10/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.11/attention/self/Transpose_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.11/attention/self/Reshape_1_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.11/attention/self/Reshape_3_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.11/attention/self/Transpose_8_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.11/attention/self/Reshape_5_output_0",
"InsertedPrecisionFreeCast_/deberta/encoder/layer.11/attention/self/Reshape_13_output_0",
"InsertedPrecisionFreeCast_/pooler/Gather_output_0"
]
},
"compile": null,
"loader": {
"task": "reranking",
"model_class": "DebertaV2ForSequenceClassification",
"model_type": "deberta-v2"
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -39,16 +39,17 @@
{
"name": "logits"
}
]
},
"optim": {
"clamp_constant_values": true
],
"compatibility": {
"transformers_attention": "eager"
}
},
"optim": {},
"quant": null,
"compile": null,
"loader": {
"task": "text-classification",
"model_class": "AutoModelForSequenceClassification",
"task": "reranking",
"model_class": "DebertaV2ForSequenceClassification",
"model_type": "deberta-v2"
}
}
}

This file was deleted.

16 changes: 16 additions & 0 deletions src/winml/modelkit/quant/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,20 @@ class WinMLQuantizationConfig:
# FP16 conversion settings (only used when mode="fp16")
fp16_keep_io_types: bool = True
fp16_op_block_list: list[str] | None = None
fp16_nodes_to_exclude: list[str] | None = None

def __post_init__(self) -> None:
"""Normalize FP16 node exclusions while preserving declaration order."""
node_exclusions: object = self.fp16_nodes_to_exclude
if node_exclusions is None:
return
if not isinstance(node_exclusions, list):
msg = "fp16_nodes_to_exclude must be a list of non-empty strings"
raise TypeError(msg)
if not all(isinstance(name, str) and name for name in node_exclusions):
msg = "fp16_nodes_to_exclude entries must be non-empty strings"
raise ValueError(msg)
self.fp16_nodes_to_exclude = list(dict.fromkeys(node_exclusions))

def to_dict(self) -> dict:
"""Convert to dictionary for serialization.
Expand Down Expand Up @@ -168,6 +182,7 @@ def to_dict(self) -> dict:
if self.mode == "fp16":
result["fp16_keep_io_types"] = self.fp16_keep_io_types
result["fp16_op_block_list"] = self.fp16_op_block_list
result["fp16_nodes_to_exclude"] = self.fp16_nodes_to_exclude
return result

@classmethod
Expand Down Expand Up @@ -217,6 +232,7 @@ def from_dict(cls, data: dict) -> WinMLQuantizationConfig:
reduce_range=data.get("reduce_range", False),
fp16_keep_io_types=data.get("fp16_keep_io_types", True),
fp16_op_block_list=data.get("fp16_op_block_list"),
fp16_nodes_to_exclude=data.get("fp16_nodes_to_exclude"),
)


Expand Down
13 changes: 11 additions & 2 deletions src/winml/modelkit/quant/fp16.py
Original file line number Diff line number Diff line change
Expand Up @@ -2822,7 +2822,7 @@ def _validate_local_function_conversion(model: ModelProto) -> None:
raise RuntimeError(msg) from error


def _validate_converted_types(model: ModelProto) -> None:
def _validate_converted_types(model: ModelProto, *, required: bool = True) -> bool:
"""Reject converted graphs whose concrete types no longer agree."""
from onnx import checker, shape_inference

Expand All @@ -2831,11 +2831,16 @@ def _validate_converted_types(model: ModelProto) -> None:
inferred = shape_inference.infer_shapes(model, check_type=True, strict_mode=True)
checker.check_model(inferred)
except (
AttributeError,
EncodeError,
checker.ValidationError,
shape_inference.InferenceError,
) as error:
if not required:
return False
msg = "FP16 conversion produced incompatible FP16 types."
raise RuntimeError(msg) from error
return True


def _node_dependencies(node: NodeProto) -> set[str]:
Expand Down Expand Up @@ -2889,6 +2894,7 @@ def convert_to_fp16(
*,
keep_io_types: bool = True,
op_block_list: list[str] | None = None,
node_block_list: list[str] | None = None,
) -> ModelProto:
"""Convert an ONNX model from FP32 to FP16 precision.

Expand Down Expand Up @@ -3024,11 +3030,13 @@ def convert_to_fp16(
if op_block_list:
logger.info(" Keeping ops in FP32: %s", op_block_list)

supports_strict_type_validation = _validate_converted_types(conversion_model, required=False)
try:
converted: ModelProto = convert_float_to_float16(
conversion_model,
keep_io_types=keep_io_types,
op_block_list=op_block_list,
node_block_list=node_block_list,
)
except EncodeError:
logger.warning(
Expand All @@ -3041,6 +3049,7 @@ def convert_to_fp16(
keep_io_types=keep_io_types,
disable_shape_infer=True,
op_block_list=op_block_list,
node_block_list=node_block_list,
)

converted_graphs = _ort_traversed_graphs(converted, op_block_list)
Expand All @@ -3060,7 +3069,7 @@ def convert_to_fp16(
_repair_sparse_float_initializers(converted, op_block_list, keep_io_types=keep_io_types)
_graph_topological_sort(converted.graph)
_validate_initializer_output_types(converted)
if requires_attribute_validation or external_attribute_tensors or requires_container_validation:
if supports_strict_type_validation:
_validate_converted_types(converted)
_validate_local_function_conversion(converted)

Expand Down
2 changes: 2 additions & 0 deletions src/winml/modelkit/quant/passes/fp16.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ class FP16Pass(BaseQuantPass):

- ``fp16_keep_io_types`` — keep model inputs/outputs in their original dtype
- ``fp16_op_block_list`` — op types that must not be cast to FP16
- ``fp16_nodes_to_exclude`` — exact node names that must not be cast to FP16

Example::

Expand Down Expand Up @@ -64,6 +65,7 @@ def run(
model,
keep_io_types=self._config.fp16_keep_io_types,
op_block_list=self._config.fp16_op_block_list,
node_block_list=self._config.fp16_nodes_to_exclude,
)
output_path.parent.mkdir(parents=True, exist_ok=True)
save_onnx(model, output_path, use_external_data=use_external_data)
Expand Down
Loading
Loading