Download test_assemble_head.py from Accio-Lab/occamy-1.0-MTP: direct link, hf CLI and curl.
- Browser
- Download file 2.9 kB
-
https://huggingface.co/Accio-Lab/occamy-1.0-MTP/resolve/main/test_assemble_head.py
- Command line
-
hf download hf://Accio-Lab/occamy-1.0-MTP/test_assemble_head.py
-
curl -L -o test_assemble_head.py https://huggingface.co/Accio-Lab/occamy-1.0-MTP/resolve/main/test_assemble_head.py
2.9 kB
| import copy | |
| import json | |
| from pathlib import Path | |
| import tempfile | |
| import unittest | |
| from unittest.mock import patch | |
| import assemble_head as assembly | |
| class AssemblyTests(unittest.TestCase): | |
| def test_nvfp4_preserves_existing_exclusions_and_is_idempotent(self): | |
| config = {'quantization_config': {'quant_algo': 'NVFP4', 'ignore': ['lm_head']}} | |
| quant = {'quantization': {'quant_algo': 'NVFP4', 'exclude_modules': ['model.visual*']}} | |
| assembly.preserve_bf16_head(config, quant) | |
| once = copy.deepcopy((config, quant)) | |
| assembly.preserve_bf16_head(config, quant) | |
| self.assertEqual(once, (config, quant)) | |
| self.assertEqual(config['quantization_config']['ignore'], ['lm_head', 'mtp.layers.0*', 'mtp*']) | |
| self.assertIn('model.visual*', quant['quantization']['exclude_modules']) | |
| def test_bf16_and_other_formats_unchanged(self): | |
| for config, quant in [({}, None), ({'quantization_config': {'quant_algo': 'FP8'}}, None)]: | |
| before = copy.deepcopy((config, quant)) | |
| assembly.preserve_bf16_head(config, quant) | |
| self.assertEqual(before, (config, quant)) | |
| def test_assembly_does_not_modify_base_or_link_quant_config(self): | |
| with tempfile.TemporaryDirectory() as root: | |
| root = Path(root) | |
| base = root/'base' | |
| base.mkdir() | |
| (base/'config.json').write_text(json.dumps({'text_config': {'model_type': 'qwen3_5_moe_text'}, 'quantization_config': {'quant_algo': 'NVFP4', 'ignore': []}})) | |
| (base/'hf_quant_config.json').write_text(json.dumps({'quantization': {'quant_algo': 'NVFP4', 'exclude_modules': []}})) | |
| (base/'model.safetensors.index.json').write_text(json.dumps({'weight_map': {'model.weight': 'base.safetensors'}})) | |
| (base/'base.safetensors').write_bytes(b'fixture') | |
| head = root/'head.safetensors' | |
| head.write_bytes(b'fixture') | |
| before = {f.name: f.read_bytes() for f in base.iterdir()} | |
| with patch.object(assembly, 'safe_open') as opened, patch('sys.argv', ['assemble', '--base', str(base), '--head', str(head), '--out', str(root/'out')]): | |
| handle = opened.return_value.__enter__.return_value | |
| handle.keys.return_value = ['mtp.weight'] | |
| handle.get_slice.return_value.get_dtype.return_value = 'BF16' | |
| handle.get_slice.return_value.get_shape.return_value = [2, 2] | |
| assembly.main() | |
| out = root/'out' | |
| self.assertFalse((out/'hf_quant_config.json').is_symlink()) | |
| self.assertIn('mtp*', json.loads((out/'hf_quant_config.json').read_text())['quantization']['exclude_modules']) | |
| self.assertTrue((out/'base.safetensors').is_symlink()) | |
| self.assertEqual(before, {f.name: f.read_bytes() for f in base.iterdir()}) | |
| if __name__ == '__main__': | |
| unittest.main() | |