File size: 3,468 Bytes
2bc6021 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 | """Regression checks for file isolation, calibration separation and task scoring."""
import json
import tempfile
import unittest
from pathlib import Path
import torch
from safetensors.torch import save_file, load_file
from checkpoint import clone
from quantize import replace_tensors
from evaluate import gsm_score, json_equal, aggregate
from code_runner import equal
from gptq_utils import accumulate_hessian
class WorkflowTests(unittest.TestCase):
def test_capped_answers_count_as_wrong_even_if_partial_answer_matches(self):
row={"input_tokens":10,"output_tokens":32768,"truncated":True,"error":None,"correct":True,"model_seconds":200}
result=aggregate([row])
self.assertEqual(result["accuracy"],0)
self.assertEqual(result["correct"],0)
self.assertEqual(result["attempted"],1)
self.assertEqual(result["truncated_counted_as_wrong"],1)
def test_hessian_remains_fp32_inside_bf16_autocast(self):
torch.manual_seed(42)
x = torch.randn(64,128)
h = torch.zeros(128,128)
with torch.autocast("cpu",dtype=torch.bfloat16):
h,n = accumulate_hessian(h,x,0)
reference = 2/len(x) * (x.T @ x)
self.assertEqual(n,64)
self.assertTrue(torch.allclose(h,reference,rtol=1e-5,atol=1e-6))
def test_quant_export_does_not_mutate_hardlinked_baseline(self):
with tempfile.TemporaryDirectory() as td:
source, target = Path(td)/"source", Path(td)/"fast"
source.mkdir()
old = torch.ones(2, 4, dtype=torch.int32)
save_file({"lm_head.weight_packed": old, "untouched.weight": torch.ones(2)}, source/"model.safetensors")
(source/"model.safetensors.index.json").write_text(json.dumps({"weight_map":{"lm_head.weight_packed":"model.safetensors"}}))
(source/"config.json").write_text(json.dumps({"quantization_config":{"config_groups":{"group_1":{"weights":{"num_bits":8}}}}}))
clone(source,target)
self.assertEqual((source/"model.safetensors").stat().st_ino,(target/"model.safetensors").stat().st_ino)
replace_tensors(target,{"lm_head":{"weight_packed":torch.zeros_like(old)}},{"group_1":4})
self.assertTrue(torch.equal(load_file(source/"model.safetensors")["lm_head.weight_packed"],old))
self.assertEqual(load_file(target/"model.safetensors")["lm_head.weight_packed"].sum().item(),0)
self.assertTrue(torch.equal(load_file(target/"model.safetensors")["untouched.weight"],torch.ones(2)))
def test_scoring_rejects_wrong_types_and_nonfinite_numbers(self):
self.assertFalse(json_equal({"x":True},{"x":1}))
self.assertFalse(json_equal(float("nan"),1))
self.assertFalse(equal("9007199254740993","9007199254740992"))
self.assertTrue(equal("1.0000001\n2","1 2"))
self.assertTrue(gsm_score("Final answer: 1,250", "1250"))
self.assertFalse(gsm_score("Final answer: 1251", "1250"))
def test_failures_remain_in_time_and_quality_denominators(self):
common={"input_tokens":10,"output_tokens":20,"truncated":False,"error":None}
result=aggregate([dict(common,correct=True,model_seconds=2),dict(common,correct=False,model_seconds=8)])
self.assertEqual(result["accuracy"],.5)
self.assertEqual(result["mean_model_seconds"],5)
self.assertEqual(result["summed_request_seconds_per_correct"],10)
if __name__ == "__main__":
unittest.main()
|