Test parametrization with different block sizes to verify flexibility.
(device, dtype, blocksize)
| 338 | [64, 128, 256], |
| 339 | ) |
| 340 | def test_different_blocksizes(device, dtype, blocksize): |
| 341 | """Test parametrization with different block sizes to verify flexibility.""" |
| 342 | if device == "hpu" and not is_supported_on_hpu("nf4", dtype): |
| 343 | pytest.skip("Configuration not supported on HPU.") |
| 344 | |
| 345 | module = ParametrizeTestModule(device=device, dtype=dtype) |
| 346 | original_param = module.expert_weights.clone() |
| 347 | |
| 348 | # Apply parametrization with specified block size |
| 349 | replace_parameter_4bit(module, "expert_weights", quant_type="nf4", blocksize=blocksize) |
| 350 | |
| 351 | # Verify reconstruction works with different block sizes |
| 352 | reconstructed = module.expert_weights |
| 353 | assert reconstructed.shape == original_param.shape, "Shape should be preserved" |
| 354 | assert reconstructed.device.type == device, "Device should match" |
| 355 | |
| 356 | # Verify quantization quality using error calculation approach from functional tests |
| 357 | err = (original_param - reconstructed.detach()).abs().float() |
| 358 | relerr = (err / (original_param.abs().float() + 1e-8)).mean() |
| 359 | err_mean = err.mean() |
| 360 | |
| 361 | # Expected (mean, std) for NF4, from 200 samples on RTX 4090. Worst-case std across dtypes. |
| 362 | N_SIGMA = 7 |
| 363 | expected_abs = {64: (0.072796, 0.000072), 128: (0.076839, 0.000093), 256: (0.080322, 0.000100)} |
| 364 | expected_rel = {64: (0.203353, 0.000326), 128: (0.215258, 0.000367), 256: (0.226056, 0.000392)} |
| 365 | |
| 366 | abs_mean, abs_std = expected_abs[blocksize] |
| 367 | rel_mean, rel_std = expected_rel[blocksize] |
| 368 | assert err_mean < abs_mean + N_SIGMA * abs_std, ( |
| 369 | f"Mean abs error {err_mean:.6f} exceeds {abs_mean:.6f} + {N_SIGMA}*{abs_std:.6f} for blocksize {blocksize}" |
| 370 | ) |
| 371 | assert relerr < rel_mean + N_SIGMA * rel_std, ( |
| 372 | f"Mean rel error {relerr:.6f} exceeds {rel_mean:.6f} + {N_SIGMA}*{rel_std:.6f} for blocksize {blocksize}" |
| 373 | ) |
| 374 | |
| 375 | |
| 376 | def test_parametrization_forward_method(): |
nothing calls this directly
no test coverage detected