|
14 | 14 | from coreai_opt._utils.export_utils import ( |
15 | 15 | clear_parametrization_original as _clear_parametrization_original, |
16 | 16 | prepare_mmap_dir as _prepare_mmap_dir, |
17 | | - validate_coreml_compatibility, |
| 17 | + validate_coreml_palettization_compatibility, |
18 | 18 | ) |
19 | 19 | from coreai_opt._utils.import_utils import lazy_import_coreai_torch |
20 | 20 | from coreai_opt._utils.metadata_utils import CompressionType, MILCompressionMetadata |
21 | 21 | from coreai_opt._utils.torch_utils import ( |
22 | 22 | mmap_module_state_dict as _mmap_module_state_dict, |
23 | 23 | ) |
24 | 24 | from coreai_opt.common import ExportBackend |
25 | | -from coreai_opt.config.spec import CompressionTargetTensor |
26 | 25 | from coreai_opt.palettization.spec.fake_palettize import ( |
27 | 26 | _FakePalettizeImplBase, |
28 | 27 | ) |
@@ -430,12 +429,16 @@ def prepare_for_mil_export(model: nn.Module) -> nn.Module: |
430 | 429 | continue |
431 | 430 | for param_name, parametrizations in module.parametrizations.items(): |
432 | 431 | _, fake_palett_mod = _find_fake_palett_parametrization(parametrizations) |
433 | | - if fake_palett_mod is not None and fake_palett_mod.lut_qspec is not None: |
434 | | - validate_coreml_compatibility( |
435 | | - CompressionTargetTensor.LUT, |
436 | | - fake_palett_mod.lut_qspec.dtype, |
437 | | - f"LUT of parameter '{param_name}' of module '{module_name}'", |
438 | | - ) |
| 432 | + if fake_palett_mod is None: |
| 433 | + continue |
| 434 | + |
| 435 | + context = f"parameter '{param_name}' of module '{module_name}'" |
| 436 | + validate_coreml_palettization_compatibility( |
| 437 | + fake_palett_mod.cluster_dim, |
| 438 | + fake_palett_mod.lut_qspec, |
| 439 | + fake_palett_mod.enable_per_channel_scale, |
| 440 | + context, |
| 441 | + ) |
439 | 442 |
|
440 | 443 | _process_weight_palettization(model, backend=ExportBackend.CoreML) |
441 | 444 |
|
|
0 commit comments