"""Check that 4-bit NF4 quantization preserves fp32 verdicts.""" import sys from pathlib import Path import torch from transformers import BitsAndBytesConfig sys.path.insert(0, str(Path(__file__).parent)) from bench_quant import load, probs # noqa: E402 THRESHOLD = 0.5 TEXTS = [ "The quick brown fox jumps over the lazy dog, or so the saying goes.", "In conclusion, the multifaceted nature of this phenomenon necessitates a " "comprehensive reevaluation of our underlying assumptions.", "My neighbour keeps parking in front of my driveway and it drives me mad.", "I'm not sure whether I left the oven on this morning before work.", "Consequently, stakeholders across the value chain must align their " "incentives to foster sustainable growth trajectories.", "We tried the new ramen place on Thursday. Broth was good, noodles okay.", "It is crucial to note that the results underscore the importance of " "systematic approaches to problem-solving in contemporary environments.", "The dog barked at the postman again, and then went back to sleep on the rug.", "By leveraging cutting-edge methodologies, we can unlock pivotal insights " "that drive transformative outcomes for all parties involved.", "Everything happens for a reason, and that reason is usually traffic.", "The document outlines several key considerations that must be addressed " "prior to implementation of the proposed framework.", "Honestly I just want to finish this report and go home.", "Through the lens of interdisciplinary inquiry, the discourse surrounding " "this topic continues to evolve in ways that defy simple characterization.", "She sold her car and bought a bike instead, which she now regrets in winter.", "This paper presents a novel approach that seamlessly integrates the " "strengths of prior work while mitigating its principal limitations.", "I forgot to water the plants again and now the basil looks tragic.", ] if __name__ == "__main__": fp32 = load() reference = probs(fp32, TEXTS) del fp32 nf4 = load( quantization_config=BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.float32, ) ) nf4.to(torch.device("mps")) quantized = probs(nf4, TEXTS) deltas = [abs(a - b) for a, b in zip(reference, quantized)] flips = [ (ref, got, text) for ref, got, text in zip(reference, quantized, TEXTS) if (ref >= THRESHOLD) != (got >= THRESHOLD) ] borderline = [ (ref, got) for ref, got in zip(reference, quantized) if 0.2 < ref < 0.8 ] print(f"texts {len(TEXTS)}") print(f"max |delta p| {max(deltas):.4f}") print(f"mean |delta p| {sum(deltas) / len(deltas):.4f}") print(f"verdict flips @0.5 {len(flips)}") print(f"fp32 borderline {len(borderline)}") for ref, got, text in flips: print(f" {ref:.3f} -> {got:.3f} {text[:70]}")