File size: 2,673 Bytes
12496fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import math
import pytest
import torch
from nexora.data import prepare
from nexora.compute import Estimate, topology, kv_cache_bytes
from nexora.evaluation import pass_at_k, wilson, percentiles, word_error_rate
from nexora.posttraining import dpo_loss, group_advantages, masked_sft_loss, rejection_sample


def record(text, **kw):
    return {"id": "x", "text": text, "source": "original", "license": "MIT", "domain": "text", **kw}


def test_pipeline_filters_and_provenance(tmp_path):
    good = "The engineering document describes reliable software with independent tests."
    held = "A hidden evaluation question asks about a particular graph algorithm and its runtime."
    rows = [record(good), record(good), record("This data has an unknown license and must not be admitted.", license="unknown"),
            record("Do not admit this private credential hf_" + "a"*30), record(held),
            record("Contact test@example.com to obtain further documentation on the system.")]
    report = prepare(rows, tmp_path, holdouts=[held])
    assert report["accepted"] == 2
    assert report["rejected"] == {"exact_duplicate": 1, "license_or_provenance_or_split": 1, "secret": 1, "contamination": 1}
    text = (tmp_path / "records.jsonl").read_text()
    assert "test@example.com" not in text and "[EMAIL]" in text


def test_compute_formulas():
    r = Estimate(120e9, 12e9, 2e12, 1024).calculate()
    assert r["weight_GB"]["bf16"] == 240
    assert r["training_flops"] == 6*12e9*2e12
    assert topology(8, 8, 2, 2, 1, 16, 8)["world"] == 64
    assert kv_cache_bytes(4, 2, 32, 256) == 262144
    with pytest.raises(ValueError):
        topology(8, 8, 8, 8, 1, 8)


def test_metrics():
    assert pass_at_k(10, 2, 1) == pytest.approx(.2)
    assert pass_at_k(10, 2, 3) == pytest.approx(1-56/120)
    assert wilson(5, 10)[0] < .5 < wilson(5, 10)[1]
    assert percentiles([1, 2, 3, 4])["p95"] == 4
    assert word_error_rate("one two three", "one four three") == pytest.approx(1/3)
    with pytest.raises(ValueError):
        pass_at_k(0, 0, 1)


def test_posttraining_gradients():
    chosen = torch.tensor([2.0, 3.0], requires_grad=True)
    loss = dpo_loss(chosen, torch.zeros(2), torch.zeros(2), torch.zeros(2))
    loss.backward()
    assert (chosen.grad < 0).all()
    assert torch.equal(group_advantages(torch.ones(2, 4)), torch.zeros(2, 4))
    logits = torch.randn(1, 3, 5, requires_grad=True)
    loss = masked_sft_loss(logits, torch.tensor([[1, 2, 3]]), torch.tensor([[False, True, True]]))
    loss.backward()
    assert torch.equal(logits.grad[0, 0], torch.zeros(5))
    assert rejection_sample(["pass", "fail", "error"], lambda s: s == "pass") == ["pass"]