Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
+107
-37
@@ -6,6 +6,7 @@ from __future__ import annotations
|
||||
import argparse
|
||||
import gc
|
||||
import hashlib
|
||||
import json
|
||||
import shutil
|
||||
import sys
|
||||
import tempfile
|
||||
@@ -73,8 +74,39 @@ def _bounded_calibration_packages(debugger_type: Any, temporary_root: Path):
|
||||
finally:
|
||||
debugger_type.predict_intermediate_outputs = original_predict
|
||||
tempfile.tempdir = previous_tempdir
|
||||
for package in temporary_root.glob("*.mlpackage"):
|
||||
shutil.rmtree(package)
|
||||
|
||||
|
||||
def _activation_cache_contract(
|
||||
source_sha256: str,
|
||||
calibration: Sequence[dict[str, Any]],
|
||||
*,
|
||||
calibration_seed: int,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"schemaVersion": 1,
|
||||
"sourcePackageSha256": source_sha256,
|
||||
"coremltoolsVersion": COREMLTOOLS_VERSION,
|
||||
"activationQuantization": "linear:per-tensor-asymmetric-uint8",
|
||||
"activationOpTypes": ["linear"],
|
||||
"calibrationRecords": len(calibration),
|
||||
"calibrationSeed": calibration_seed,
|
||||
"calibrationPromptHashes": sorted(
|
||||
prompt_hash(item["prompt"]) for item in calibration
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _activation_cache_is_valid(
|
||||
package: Path,
|
||||
manifest: Path,
|
||||
expected: dict[str, Any],
|
||||
) -> bool:
|
||||
if not package.is_dir() or not manifest.is_file():
|
||||
return False
|
||||
try:
|
||||
return json.loads(manifest.read_text(encoding="utf-8")) == expected
|
||||
except (OSError, UnicodeError, json.JSONDecodeError):
|
||||
return False
|
||||
|
||||
|
||||
def _optimization_configs(optimize: Any) -> tuple[Any, Any]:
|
||||
@@ -160,24 +192,38 @@ def quantize(args: argparse.Namespace) -> dict[str, Any]:
|
||||
args.calibration_records,
|
||||
seed=args.calibration_seed,
|
||||
)
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.model_dir,
|
||||
local_files_only=True,
|
||||
source_sha256 = _tree_sha256(args.model)
|
||||
activation_cache = args.output.with_name(
|
||||
f"{args.output.stem}-a8-cache.mlpackage"
|
||||
)
|
||||
activation_cache_manifest = activation_cache.with_name(
|
||||
f"{activation_cache.stem}-manifest.json"
|
||||
)
|
||||
cache_contract = _activation_cache_contract(
|
||||
source_sha256,
|
||||
calibration,
|
||||
calibration_seed=args.calibration_seed,
|
||||
)
|
||||
reuse_activation_cache = _activation_cache_is_valid(
|
||||
activation_cache,
|
||||
activation_cache_manifest,
|
||||
cache_contract,
|
||||
)
|
||||
sample_data = []
|
||||
for record in calibration:
|
||||
encoded = encode_fixed_shape(tokenizer, [record["prompt"]], torch)
|
||||
sample_data.append(
|
||||
{
|
||||
name: value.numpy().astype(np.int32, copy=False)
|
||||
for name, value in encoded.items()
|
||||
}
|
||||
if not reuse_activation_cache:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.model_dir,
|
||||
local_files_only=True,
|
||||
)
|
||||
for record in calibration:
|
||||
encoded = encode_fixed_shape(tokenizer, [record["prompt"]], torch)
|
||||
sample_data.append(
|
||||
{
|
||||
name: value.numpy().astype(np.int32, copy=False)
|
||||
for name, value in encoded.items()
|
||||
}
|
||||
)
|
||||
|
||||
print(
|
||||
f"Core ML activation calibration: {len(sample_data)} records",
|
||||
flush=True,
|
||||
)
|
||||
source_model = ct.models.MLModel(
|
||||
str(args.model),
|
||||
compute_units=ct.ComputeUnit.CPU_ONLY,
|
||||
@@ -196,34 +242,58 @@ def quantize(args: argparse.Namespace) -> dict[str, Any]:
|
||||
dir=args.output.parent,
|
||||
) as temporary:
|
||||
with _bounded_calibration_packages(ModelDebugger, Path(temporary)):
|
||||
activation_quantized = cto.coreml.linear_quantize_activations(
|
||||
source_model,
|
||||
activation_config,
|
||||
sample_data,
|
||||
calibration_op_group_size=args.calibration_op_group_size,
|
||||
if reuse_activation_cache:
|
||||
print(
|
||||
f"Reusing Core ML A8 cache: {activation_cache}",
|
||||
flush=True,
|
||||
)
|
||||
activation_quantized = ct.models.MLModel(
|
||||
str(activation_cache),
|
||||
compute_units=ct.ComputeUnit.CPU_ONLY,
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"Core ML activation calibration: {len(sample_data)} records",
|
||||
flush=True,
|
||||
)
|
||||
activation_quantized = cto.coreml.linear_quantize_activations(
|
||||
source_model,
|
||||
activation_config,
|
||||
sample_data,
|
||||
calibration_op_group_size=args.calibration_op_group_size,
|
||||
)
|
||||
activation_quantized.save(str(activation_cache))
|
||||
write_json(activation_cache_manifest, cache_contract)
|
||||
print(
|
||||
f"Core ML A8 cache: {activation_cache}",
|
||||
flush=True,
|
||||
)
|
||||
print("Core ML weight quantization: W8", flush=True)
|
||||
quantized = cto.coreml.linear_quantize_weights(
|
||||
activation_quantized,
|
||||
weight_config,
|
||||
)
|
||||
print("Core ML weight quantization: W8", flush=True)
|
||||
quantized = cto.coreml.linear_quantize_weights(
|
||||
activation_quantized,
|
||||
weight_config,
|
||||
)
|
||||
quantized.user_defined_metadata["com.nucleic.model.quantization"] = "W8A8"
|
||||
quantized.user_defined_metadata["com.nucleic.model.quantizationCalibration"] = (
|
||||
f"stratified:{args.calibration_records}:seed={args.calibration_seed}"
|
||||
)
|
||||
quantized.user_defined_metadata["com.nucleic.model.quantization"] = (
|
||||
"W8A8"
|
||||
)
|
||||
quantized.user_defined_metadata[
|
||||
"com.nucleic.model.quantizationCalibration"
|
||||
] = f"stratified:{len(calibration)}:seed={args.calibration_seed}"
|
||||
|
||||
if args.output.exists():
|
||||
if args.output.is_dir():
|
||||
shutil.rmtree(args.output)
|
||||
else:
|
||||
args.output.unlink()
|
||||
quantized.save(str(args.output))
|
||||
if args.output.exists():
|
||||
if args.output.is_dir():
|
||||
shutil.rmtree(args.output)
|
||||
else:
|
||||
args.output.unlink()
|
||||
# Save while the Core ML Tools result package and its source weights are
|
||||
# still alive inside the dedicated temporary directory.
|
||||
quantized.save(str(args.output))
|
||||
|
||||
manifest = _package_manifest(args.output)
|
||||
manifest.update(
|
||||
{
|
||||
"sourcePackage": str(args.model),
|
||||
"sourcePackageSha256": _tree_sha256(args.model),
|
||||
"sourcePackageSha256": source_sha256,
|
||||
"coremltoolsVersion": ct.__version__,
|
||||
"quantization": {
|
||||
"name": "W8A8",
|
||||
|
||||
Reference in New Issue
Block a user