mirror of
https://github.com/priyanshujain/sanderling.git
synced 2026-10-02 19:17:10 +00:00
feat(analyze): survival analysis over campaign directories
Steps to first violation with clean runs right-censored at the budget, since per-run yield is a binary at 11 to 45 percent and separating two arms on it would need roughly 80 runs per arm. Kaplan-Meier, log-rank, Wilcoxon rank-sum with Vargha-Delaney A12, Holm within each family. A hand-rolled log-rank that is subtly wrong is a silent-wrong-number generator and would be believed, so every statistic is validated against a published worked example with the source named in the test: R survdiff on aml, Freireich 6-MP, Hollander and Wolfe 1973 for the rank sum, printed p.adjust output for Holm. Two could not be: the k>2 log-rank, guarded by calibration instead, and the tie-corrected variance, checked against an exact permutation variance. Failed and timed-out runs are excluded as missing data and counted by reason, never treated as censored observations, which would bias the result. Claude-Session: https://claude.ai/code/session_01A5KmftdEJ49A9z5mF5ESrX
This commit is contained in:
1 parent
71dffef2f2
commit
019d608f65
16 files changed
+2595
No files matched your search
@@ -0,0 +1,192 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type armSummary struct {
|
||||||
|
Arm string `json:"arm"`
|
||||||
|
Generator string `json:"generator,omitempty"`
|
||||||
|
Platform string `json:"platform,omitempty"`
|
||||||
|
StepBudget int `json:"step_budget"`
|
||||||
|
Directories []string `json:"directories"`
|
||||||
|
Recorded int `json:"recorded_runs"`
|
||||||
|
Usable int `json:"usable_runs"`
|
||||||
|
Violated int `json:"violated_runs"`
|
||||||
|
Censored int `json:"censored_runs"`
|
||||||
|
Excluded int `json:"excluded_runs"`
|
||||||
|
ExcludedByReason map[string]int `json:"excluded_by_reason,omitempty"`
|
||||||
|
MissingSeeds []int64 `json:"missing_seeds,omitempty"`
|
||||||
|
EventsHeldAtBudget int `json:"events_held_at_budget"`
|
||||||
|
MedianStepsToFirstViolation *float64 `json:"median_steps_to_first_violation"`
|
||||||
|
SurvivalCurve []survivalPoint `json:"survival_curve,omitempty"`
|
||||||
|
ViolationRate *float64 `json:"violation_rate"`
|
||||||
|
TotalActions int `json:"total_actions"`
|
||||||
|
TotalRunHours float64 `json:"total_run_hours"`
|
||||||
|
Detections int `json:"detections"`
|
||||||
|
DefectsPerThousandActions *float64 `json:"defects_per_thousand_actions"`
|
||||||
|
DefectsPerHour *float64 `json:"defects_per_hour"`
|
||||||
|
DistinctDefects int `json:"distinct_defects"`
|
||||||
|
SingletonDefects int `json:"singleton_defects"`
|
||||||
|
SingletonFraction *float64 `json:"singleton_fraction"`
|
||||||
|
DefectRunCounts map[string]int `json:"defect_run_counts,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type pairwiseResult struct {
|
||||||
|
First string `json:"first"`
|
||||||
|
Second string `json:"second"`
|
||||||
|
FirstSize int `json:"first_size"`
|
||||||
|
SecondSize int `json:"second_size"`
|
||||||
|
Statistic float64 `json:"mann_whitney_u"`
|
||||||
|
A12 float64 `json:"a12"`
|
||||||
|
PValue float64 `json:"p_value"`
|
||||||
|
HolmPValue float64 `json:"holm_p_value"`
|
||||||
|
Exact bool `json:"exact"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type analysis struct {
|
||||||
|
GeneratedAt time.Time `json:"generated_at"`
|
||||||
|
Outcome string `json:"outcome"`
|
||||||
|
Arms []armSummary `json:"arms"`
|
||||||
|
LogRank *logRankResult `json:"log_rank"`
|
||||||
|
Pairwise []pairwiseResult `json:"pairwise"`
|
||||||
|
Notes []string `json:"notes,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
const outcomeDescription = "steps to first violation, right-censored at the step budget"
|
||||||
|
|
||||||
|
func analyse(arms []arm, now time.Time) analysis {
|
||||||
|
result := analysis{GeneratedAt: now, Outcome: outcomeDescription}
|
||||||
|
for _, current := range arms {
|
||||||
|
result.Arms = append(result.Arms, summarize(current))
|
||||||
|
}
|
||||||
|
|
||||||
|
var testable []arm
|
||||||
|
for _, current := range arms {
|
||||||
|
if len(current.observations()) > 0 {
|
||||||
|
testable = append(testable, current)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(testable) < len(arms) {
|
||||||
|
result.Notes = append(result.Notes,
|
||||||
|
"arms with no usable runs are reported but left out of the log-rank test and the pairwise comparisons")
|
||||||
|
}
|
||||||
|
if len(testable) >= 2 {
|
||||||
|
names := make([]string, len(testable))
|
||||||
|
groups := make([][]observation, len(testable))
|
||||||
|
for index, current := range testable {
|
||||||
|
names[index] = current.Name
|
||||||
|
groups[index] = current.observations()
|
||||||
|
}
|
||||||
|
test := logRank(names, groups)
|
||||||
|
result.LogRank = &test
|
||||||
|
result.Pairwise = comparePairs(testable)
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func comparePairs(arms []arm) []pairwiseResult {
|
||||||
|
var pairs []pairwiseResult
|
||||||
|
for first := 0; first < len(arms); first++ {
|
||||||
|
for second := first + 1; second < len(arms); second++ {
|
||||||
|
test := rankSum(arms[first].stepTimes(), arms[second].stepTimes())
|
||||||
|
pairs = append(pairs, pairwiseResult{
|
||||||
|
First: arms[first].Name,
|
||||||
|
Second: arms[second].Name,
|
||||||
|
FirstSize: test.FirstSize,
|
||||||
|
SecondSize: test.SecondSize,
|
||||||
|
Statistic: test.Statistic,
|
||||||
|
A12: test.A12,
|
||||||
|
PValue: test.PValue,
|
||||||
|
HolmPValue: math.NaN(),
|
||||||
|
Exact: test.Exact,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Holm runs over this one family of comparisons. A comparison whose p-value
|
||||||
|
// could not be computed is not part of the family and does not shrink the
|
||||||
|
// correction the others receive.
|
||||||
|
var family []int
|
||||||
|
var raw []float64
|
||||||
|
for index, pair := range pairs {
|
||||||
|
if math.IsNaN(pair.PValue) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
family = append(family, index)
|
||||||
|
raw = append(raw, pair.PValue)
|
||||||
|
}
|
||||||
|
for position, adjusted := range holm(raw) {
|
||||||
|
pairs[family[position]].HolmPValue = adjusted
|
||||||
|
}
|
||||||
|
return pairs
|
||||||
|
}
|
||||||
|
|
||||||
|
func summarize(current arm) armSummary {
|
||||||
|
summary := armSummary{
|
||||||
|
Arm: current.Name,
|
||||||
|
Generator: current.Generator,
|
||||||
|
Platform: current.Platform,
|
||||||
|
StepBudget: current.Budget,
|
||||||
|
Directories: current.Directories,
|
||||||
|
Recorded: len(current.Runs),
|
||||||
|
MissingSeeds: current.MissingSeeds,
|
||||||
|
}
|
||||||
|
runsPerDefect := map[string]int{}
|
||||||
|
for _, item := range current.Runs {
|
||||||
|
if item.ExcludedBecause != "" {
|
||||||
|
summary.Excluded++
|
||||||
|
if summary.ExcludedByReason == nil {
|
||||||
|
summary.ExcludedByReason = map[string]int{}
|
||||||
|
}
|
||||||
|
summary.ExcludedByReason[item.ExcludedBecause]++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
summary.Usable++
|
||||||
|
summary.TotalActions += item.Steps
|
||||||
|
summary.TotalRunHours += float64(item.DurationMillis) / float64(time.Hour/time.Millisecond)
|
||||||
|
if item.ClampedToBudget {
|
||||||
|
summary.EventsHeldAtBudget++
|
||||||
|
}
|
||||||
|
if item.Violated {
|
||||||
|
summary.Violated++
|
||||||
|
} else {
|
||||||
|
summary.Censored++
|
||||||
|
}
|
||||||
|
distinct := slices.Compact(slices.Sorted(slices.Values(item.ViolatedProperties)))
|
||||||
|
summary.Detections += len(distinct)
|
||||||
|
for _, property := range distinct {
|
||||||
|
runsPerDefect[property]++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
summary.SurvivalCurve = kaplanMeier(current.observations())
|
||||||
|
if median, ok := medianSurvival(summary.SurvivalCurve); ok {
|
||||||
|
summary.MedianStepsToFirstViolation = &median
|
||||||
|
}
|
||||||
|
if summary.Usable > 0 {
|
||||||
|
rate := float64(summary.Violated) / float64(summary.Usable)
|
||||||
|
summary.ViolationRate = &rate
|
||||||
|
}
|
||||||
|
if summary.TotalActions > 0 {
|
||||||
|
perThousand := 1000 * float64(summary.Detections) / float64(summary.TotalActions)
|
||||||
|
summary.DefectsPerThousandActions = &perThousand
|
||||||
|
}
|
||||||
|
if summary.TotalRunHours > 0 {
|
||||||
|
perHour := float64(summary.Detections) / summary.TotalRunHours
|
||||||
|
summary.DefectsPerHour = &perHour
|
||||||
|
}
|
||||||
|
if len(runsPerDefect) > 0 {
|
||||||
|
summary.DefectRunCounts = runsPerDefect
|
||||||
|
summary.DistinctDefects = len(runsPerDefect)
|
||||||
|
for _, count := range runsPerDefect {
|
||||||
|
if count == 1 {
|
||||||
|
summary.SingletonDefects++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fraction := float64(summary.SingletonDefects) / float64(summary.DistinctDefects)
|
||||||
|
summary.SingletonFraction = &fraction
|
||||||
|
}
|
||||||
|
return summary
|
||||||
|
}
|
||||||
@@ -0,0 +1,218 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func violatingRun(seed int64, steps, origin int, properties ...string) classifiedRun {
|
||||||
|
return classifiedRun{
|
||||||
|
Seed: seed,
|
||||||
|
Steps: steps,
|
||||||
|
DurationMillis: 60_000,
|
||||||
|
OriginStep: origin,
|
||||||
|
Violated: true,
|
||||||
|
ViolatedProperties: properties,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cleanRun(seed int64, steps int) classifiedRun {
|
||||||
|
return classifiedRun{Seed: seed, Steps: steps, DurationMillis: 60_000}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSummarize_ArmWhereNoRunViolated(t *testing.T) {
|
||||||
|
summary := summarize(arm{
|
||||||
|
Name: "quiet",
|
||||||
|
Budget: 40,
|
||||||
|
Runs: []classifiedRun{cleanRun(1, 40), cleanRun(2, 40), cleanRun(3, 40)},
|
||||||
|
})
|
||||||
|
if summary.Usable != 3 || summary.Censored != 3 || summary.Violated != 0 {
|
||||||
|
t.Errorf("summary %+v, want three censored runs", summary)
|
||||||
|
}
|
||||||
|
if summary.MedianStepsToFirstViolation != nil {
|
||||||
|
t.Errorf("median %v, want undefined", *summary.MedianStepsToFirstViolation)
|
||||||
|
}
|
||||||
|
if summary.ViolationRate == nil || *summary.ViolationRate != 0 {
|
||||||
|
t.Errorf("violation rate %v, want 0", summary.ViolationRate)
|
||||||
|
}
|
||||||
|
if summary.DistinctDefects != 0 || summary.SingletonFraction != nil {
|
||||||
|
t.Errorf("defects %d singleton fraction %v, want none", summary.DistinctDefects, summary.SingletonFraction)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSummarize_ArmWhereEveryRunViolated(t *testing.T) {
|
||||||
|
summary := summarize(arm{
|
||||||
|
Name: "loud",
|
||||||
|
Budget: 40,
|
||||||
|
Runs: []classifiedRun{
|
||||||
|
violatingRun(1, 5, 5, "cartTotal"),
|
||||||
|
violatingRun(2, 9, 9, "cartTotal"),
|
||||||
|
violatingRun(3, 11, 11, "cartTotal", "backNavigation"),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if summary.Violated != 3 || summary.Censored != 0 {
|
||||||
|
t.Errorf("summary %+v, want three events", summary)
|
||||||
|
}
|
||||||
|
if summary.MedianStepsToFirstViolation == nil || *summary.MedianStepsToFirstViolation != 9 {
|
||||||
|
t.Errorf("median %v, want 9", summary.MedianStepsToFirstViolation)
|
||||||
|
}
|
||||||
|
if *summary.ViolationRate != 1 {
|
||||||
|
t.Errorf("violation rate %v, want 1", *summary.ViolationRate)
|
||||||
|
}
|
||||||
|
if summary.Detections != 4 || summary.DistinctDefects != 2 {
|
||||||
|
t.Errorf("detections %d distinct %d, want 4 and 2", summary.Detections, summary.DistinctDefects)
|
||||||
|
}
|
||||||
|
// backNavigation appears in one run of three, cartTotal in all three.
|
||||||
|
if summary.SingletonDefects != 1 || math.Abs(*summary.SingletonFraction-0.5) > 1e-12 {
|
||||||
|
t.Errorf("singletons %d fraction %v, want 1 and 0.5", summary.SingletonDefects, summary.SingletonFraction)
|
||||||
|
}
|
||||||
|
// 25 actions over 3 minutes.
|
||||||
|
if math.Abs(*summary.DefectsPerThousandActions-160) > 1e-9 {
|
||||||
|
t.Errorf("defects per thousand actions %v, want 160", *summary.DefectsPerThousandActions)
|
||||||
|
}
|
||||||
|
if math.Abs(*summary.DefectsPerHour-80) > 1e-9 {
|
||||||
|
t.Errorf("defects per hour %v, want 80", *summary.DefectsPerHour)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSummarize_ArmWithNoUsableRunsAfterExclusions(t *testing.T) {
|
||||||
|
summary := summarize(arm{
|
||||||
|
Name: "broken",
|
||||||
|
Budget: 40,
|
||||||
|
Runs: []classifiedRun{
|
||||||
|
{Seed: 1, ExcludedBecause: reasonTimedOut},
|
||||||
|
{Seed: 2, ExcludedBecause: reasonNonzeroExit},
|
||||||
|
{Seed: 3, ExcludedBecause: reasonNonzeroExit},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if summary.Usable != 0 || summary.Excluded != 3 {
|
||||||
|
t.Errorf("summary %+v, want no usable runs and three exclusions", summary)
|
||||||
|
}
|
||||||
|
if summary.ExcludedByReason[reasonNonzeroExit] != 2 || summary.ExcludedByReason[reasonTimedOut] != 1 {
|
||||||
|
t.Errorf("exclusions %v", summary.ExcludedByReason)
|
||||||
|
}
|
||||||
|
if summary.ViolationRate != nil || summary.MedianStepsToFirstViolation != nil {
|
||||||
|
t.Error("reported a rate or a median for an arm with nothing in it")
|
||||||
|
}
|
||||||
|
if summary.DefectsPerThousandActions != nil || summary.DefectsPerHour != nil {
|
||||||
|
t.Error("reported a yield rate with no actions and no time")
|
||||||
|
}
|
||||||
|
if len(summary.SurvivalCurve) != 0 {
|
||||||
|
t.Errorf("survival curve %v, want empty", summary.SurvivalCurve)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSummarize_SingleRunArm(t *testing.T) {
|
||||||
|
summary := summarize(arm{Name: "one", Budget: 40, Runs: []classifiedRun{violatingRun(1, 6, 6, "cartTotal")}})
|
||||||
|
if summary.Usable != 1 || summary.Violated != 1 {
|
||||||
|
t.Errorf("summary %+v", summary)
|
||||||
|
}
|
||||||
|
if summary.MedianStepsToFirstViolation == nil || *summary.MedianStepsToFirstViolation != 6 {
|
||||||
|
t.Errorf("median %v, want 6", summary.MedianStepsToFirstViolation)
|
||||||
|
}
|
||||||
|
if summary.SingletonDefects != 1 || *summary.SingletonFraction != 1 {
|
||||||
|
t.Errorf("singletons %d fraction %v, want 1 and 1", summary.SingletonDefects, summary.SingletonFraction)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Excluded runs must not reach the survival data at all, and the counts must
|
||||||
|
// keep them visible.
|
||||||
|
func TestAnalyse_ExcludedRunsNeverBecomeObservations(t *testing.T) {
|
||||||
|
current := arm{
|
||||||
|
Name: "mixed",
|
||||||
|
Budget: 30,
|
||||||
|
Runs: []classifiedRun{
|
||||||
|
violatingRun(1, 8, 8, "cartTotal"),
|
||||||
|
cleanRun(2, 30),
|
||||||
|
{Seed: 3, ExcludedBecause: reasonTimedOut},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
observations := current.observations()
|
||||||
|
if len(observations) != 2 {
|
||||||
|
t.Fatalf("%d observations, want 2", len(observations))
|
||||||
|
}
|
||||||
|
summary := summarize(current)
|
||||||
|
if summary.Usable != 2 || summary.Excluded != 1 || summary.Violated != 1 || summary.Censored != 1 {
|
||||||
|
t.Errorf("summary %+v", summary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAnalyse_ArmWithNoUsableRunsIsReportedButNotTested(t *testing.T) {
|
||||||
|
result := analyse([]arm{
|
||||||
|
{Name: "a", Budget: 30, Runs: []classifiedRun{violatingRun(1, 4, 4), violatingRun(2, 6, 6)}},
|
||||||
|
{Name: "b", Budget: 30, Runs: []classifiedRun{cleanRun(1, 30), cleanRun(2, 30)}},
|
||||||
|
{Name: "c", Budget: 30, Runs: []classifiedRun{{Seed: 1, ExcludedBecause: reasonNonzeroExit}}},
|
||||||
|
}, time.Unix(0, 0).UTC())
|
||||||
|
|
||||||
|
if len(result.Arms) != 3 {
|
||||||
|
t.Fatalf("%d arms reported, want all 3", len(result.Arms))
|
||||||
|
}
|
||||||
|
if result.LogRank == nil || len(result.LogRank.Groups) != 2 {
|
||||||
|
t.Fatalf("log-rank %+v, want the two testable arms", result.LogRank)
|
||||||
|
}
|
||||||
|
if len(result.Pairwise) != 1 {
|
||||||
|
t.Fatalf("%d comparisons, want 1", len(result.Pairwise))
|
||||||
|
}
|
||||||
|
if len(result.Notes) == 0 {
|
||||||
|
t.Error("no note explaining the dropped arm")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// With a single testable arm there is nothing to compare against, and the tool
|
||||||
|
// must say so instead of producing a statistic.
|
||||||
|
func TestAnalyse_SingleArmHasNoTests(t *testing.T) {
|
||||||
|
result := analyse([]arm{
|
||||||
|
{Name: "a", Budget: 30, Runs: []classifiedRun{violatingRun(1, 4, 4)}},
|
||||||
|
}, time.Unix(0, 0).UTC())
|
||||||
|
if result.LogRank != nil || len(result.Pairwise) != 0 {
|
||||||
|
t.Errorf("log-rank %+v pairwise %v, want neither", result.LogRank, result.Pairwise)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Holm is applied within the family of pairwise comparisons, so with three arms
|
||||||
|
// the smallest raw p-value is multiplied by three.
|
||||||
|
func TestComparePairs_AppliesHolmWithinTheFamily(t *testing.T) {
|
||||||
|
arms := []arm{
|
||||||
|
{Name: "a", Budget: 40, Runs: manyRuns(12, 4)},
|
||||||
|
{Name: "b", Budget: 40, Runs: manyRuns(12, 20)},
|
||||||
|
{Name: "c", Budget: 40, Runs: manyRuns(12, 36)},
|
||||||
|
}
|
||||||
|
pairs := comparePairs(arms)
|
||||||
|
if len(pairs) != 3 {
|
||||||
|
t.Fatalf("%d comparisons, want 3", len(pairs))
|
||||||
|
}
|
||||||
|
raw := make([]float64, len(pairs))
|
||||||
|
for index, pair := range pairs {
|
||||||
|
raw[index] = pair.PValue
|
||||||
|
if pair.HolmPValue < pair.PValue-1e-12 {
|
||||||
|
t.Errorf("%s vs %s: holm p %v below raw p %v", pair.First, pair.Second, pair.HolmPValue, pair.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
expected := holm(raw)
|
||||||
|
for index, pair := range pairs {
|
||||||
|
if math.Abs(pair.HolmPValue-expected[index]) > 1e-12 {
|
||||||
|
t.Errorf("comparison %d holm p %v, want %v", index, pair.HolmPValue, expected[index])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// a12 above one half means the first arm needed more steps before its first
|
||||||
|
// violation, so the arm that finds defects sooner sits below one half.
|
||||||
|
func TestComparePairs_A12DirectionFollowsStepCounts(t *testing.T) {
|
||||||
|
pairs := comparePairs([]arm{
|
||||||
|
{Name: "slow", Budget: 40, Runs: manyRuns(6, 30)},
|
||||||
|
{Name: "fast", Budget: 40, Runs: manyRuns(6, 4)},
|
||||||
|
})
|
||||||
|
if pairs[0].A12 <= 0.5 {
|
||||||
|
t.Errorf("a12 %v for the slower arm listed first, want above 0.5", pairs[0].A12)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func manyRuns(count, originStep int) []classifiedRun {
|
||||||
|
runs := make([]classifiedRun, 0, count)
|
||||||
|
for index := 0; index < count; index++ {
|
||||||
|
runs = append(runs, violatingRun(int64(index), originStep+index, originStep+index, "cartTotal"))
|
||||||
|
}
|
||||||
|
return runs
|
||||||
|
}
|
||||||
@@ -0,0 +1,83 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import "math"
|
||||||
|
|
||||||
|
// standardNormalUpperTail is P(Z > z) for a standard normal Z.
|
||||||
|
func standardNormalUpperTail(z float64) float64 {
|
||||||
|
return 0.5 * math.Erfc(z/math.Sqrt2)
|
||||||
|
}
|
||||||
|
|
||||||
|
// chiSquareUpperTail is P(X > x) for a chi-square variate with the given
|
||||||
|
// degrees of freedom, which is the regularized upper incomplete gamma
|
||||||
|
// Q(degreesOfFreedom/2, x/2).
|
||||||
|
func chiSquareUpperTail(x float64, degreesOfFreedom int) float64 {
|
||||||
|
if degreesOfFreedom <= 0 || math.IsNaN(x) {
|
||||||
|
return math.NaN()
|
||||||
|
}
|
||||||
|
if x <= 0 {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return regularizedUpperGamma(float64(degreesOfFreedom)/2, x/2)
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
gammaIterationLimit = 2000
|
||||||
|
gammaTolerance = 1e-15
|
||||||
|
gammaTiny = 1e-300
|
||||||
|
)
|
||||||
|
|
||||||
|
// regularizedUpperGamma is Q(shape, x). The series is used below the crossover
|
||||||
|
// and the continued fraction above it, as in Numerical Recipes in C, 2nd ed.,
|
||||||
|
// section 6.2 (gammp/gammq).
|
||||||
|
func regularizedUpperGamma(shape, x float64) float64 {
|
||||||
|
if x < shape+1 {
|
||||||
|
return 1 - lowerGammaSeries(shape, x)
|
||||||
|
}
|
||||||
|
return upperGammaContinuedFraction(shape, x)
|
||||||
|
}
|
||||||
|
|
||||||
|
func lowerGammaSeries(shape, x float64) float64 {
|
||||||
|
term := 1 / shape
|
||||||
|
sum := term
|
||||||
|
for iteration := 1; iteration < gammaIterationLimit; iteration++ {
|
||||||
|
term *= x / (shape + float64(iteration))
|
||||||
|
sum += term
|
||||||
|
if math.Abs(term) < math.Abs(sum)*gammaTolerance {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return sum * math.Exp(-x+shape*math.Log(x)-logGamma(shape))
|
||||||
|
}
|
||||||
|
|
||||||
|
// upperGammaContinuedFraction evaluates Q(shape, x) with the modified Lentz
|
||||||
|
// algorithm, Numerical Recipes in C, 2nd ed., section 5.2.
|
||||||
|
func upperGammaContinuedFraction(shape, x float64) float64 {
|
||||||
|
b := x + 1 - shape
|
||||||
|
c := 1 / gammaTiny
|
||||||
|
d := 1 / b
|
||||||
|
h := d
|
||||||
|
for iteration := 1; iteration < gammaIterationLimit; iteration++ {
|
||||||
|
numerator := -float64(iteration) * (float64(iteration) - shape)
|
||||||
|
b += 2
|
||||||
|
d = numerator*d + b
|
||||||
|
if math.Abs(d) < gammaTiny {
|
||||||
|
d = gammaTiny
|
||||||
|
}
|
||||||
|
c = b + numerator/c
|
||||||
|
if math.Abs(c) < gammaTiny {
|
||||||
|
c = gammaTiny
|
||||||
|
}
|
||||||
|
d = 1 / d
|
||||||
|
delta := d * c
|
||||||
|
h *= delta
|
||||||
|
if math.Abs(delta-1) < gammaTolerance {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return h * math.Exp(-x+shape*math.Log(x)-logGamma(shape))
|
||||||
|
}
|
||||||
|
|
||||||
|
func logGamma(x float64) float64 {
|
||||||
|
value, _ := math.Lgamma(x)
|
||||||
|
return value
|
||||||
|
}
|
||||||
@@ -0,0 +1,74 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Chi-square critical values are the standard published table entries: the
|
||||||
|
// upper-tail probability of each of these statistics is the stated alpha in any
|
||||||
|
// chi-square table, for example Pearson and Hartley, Biometrika Tables for
|
||||||
|
// Statisticians, Table 8.
|
||||||
|
func TestChiSquareUpperTail_MatchesPublishedCriticalValues(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
statistic float64
|
||||||
|
degreesOfFreedom int
|
||||||
|
expected float64
|
||||||
|
}{
|
||||||
|
{3.841459, 1, 0.05},
|
||||||
|
{6.634897, 1, 0.01},
|
||||||
|
{10.827566, 1, 0.001},
|
||||||
|
{5.991465, 2, 0.05},
|
||||||
|
{9.210340, 2, 0.01},
|
||||||
|
{7.814728, 3, 0.05},
|
||||||
|
{11.344867, 3, 0.01},
|
||||||
|
{9.487729, 4, 0.05},
|
||||||
|
{18.307038, 10, 0.05},
|
||||||
|
}
|
||||||
|
for _, test := range cases {
|
||||||
|
got := chiSquareUpperTail(test.statistic, test.degreesOfFreedom)
|
||||||
|
if math.Abs(got-test.expected) > 1e-6 {
|
||||||
|
t.Errorf("chiSquareUpperTail(%v, %d) = %v, want %v", test.statistic, test.degreesOfFreedom, got, test.expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// For one degree of freedom the upper tail has the closed form erfc(sqrt(x/2)),
|
||||||
|
// which is an independent check on the incomplete gamma routine.
|
||||||
|
func TestChiSquareUpperTail_AgreesWithClosedFormAtOneDegreeOfFreedom(t *testing.T) {
|
||||||
|
for _, statistic := range []float64{0.1, 1, 3.4, 16.79, 40, 120} {
|
||||||
|
expected := math.Erfc(math.Sqrt(statistic / 2))
|
||||||
|
got := chiSquareUpperTail(statistic, 1)
|
||||||
|
if math.Abs(got-expected) > 1e-12*math.Max(1, expected) {
|
||||||
|
t.Errorf("chiSquareUpperTail(%v, 1) = %v, want %v", statistic, got, expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestChiSquareUpperTail_ZeroStatisticIsCertain(t *testing.T) {
|
||||||
|
if got := chiSquareUpperTail(0, 1); got != 1 {
|
||||||
|
t.Errorf("chiSquareUpperTail(0, 1) = %v, want 1", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Standard normal quantiles from any published normal table.
|
||||||
|
func TestStandardNormalUpperTail_MatchesPublishedQuantiles(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
z float64
|
||||||
|
expected float64
|
||||||
|
}{
|
||||||
|
{1.281552, 0.10},
|
||||||
|
{1.644854, 0.05},
|
||||||
|
{1.959964, 0.025},
|
||||||
|
{2.326348, 0.01},
|
||||||
|
{2.575829, 0.005},
|
||||||
|
{3.090232, 0.001},
|
||||||
|
{0, 0.5},
|
||||||
|
}
|
||||||
|
for _, test := range cases {
|
||||||
|
got := standardNormalUpperTail(test.z)
|
||||||
|
if math.Abs(got-test.expected) > 1e-6 {
|
||||||
|
t.Errorf("standardNormalUpperTail(%v) = %v, want %v", test.z, got, test.expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,246 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"math"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// buildFixtureCampaign writes a campaign directory shaped exactly like the one
|
||||||
|
// the campaign tool emits: campaign.json plus one runs.jsonl line per seed.
|
||||||
|
func buildFixtureCampaign(t *testing.T, directory, armName string, budget int, records []map[string]any) {
|
||||||
|
t.Helper()
|
||||||
|
seeds := make([]int, 0, len(records))
|
||||||
|
for _, record := range records {
|
||||||
|
seeds = append(seeds, record["seed"].(int))
|
||||||
|
}
|
||||||
|
writeCampaign(t, directory, map[string]any{
|
||||||
|
"arm": armName,
|
||||||
|
"generator": "seeded",
|
||||||
|
"platform": "web",
|
||||||
|
"spec_path": "/specs/folio.ts",
|
||||||
|
"bundle_id": "app.folio",
|
||||||
|
"max_steps": budget,
|
||||||
|
"seeds": seeds,
|
||||||
|
"host": "experiment-host",
|
||||||
|
"started_at": "2026-08-12T00:00:00Z",
|
||||||
|
"argument_temp": nil,
|
||||||
|
}, records)
|
||||||
|
}
|
||||||
|
|
||||||
|
func seededArmRecords() []map[string]any {
|
||||||
|
// Ten runs: two violate early, one violates late, six run the budget clean,
|
||||||
|
// one times out and is missing data rather than a censored observation.
|
||||||
|
return []map[string]any{
|
||||||
|
{"seed": 1, "exit_code": 0, "steps": 60, "duration_millis": 300000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 2, "exit_code": 0, "steps": 18, "duration_millis": 120000,
|
||||||
|
"first_violation_origin_step": 14, "violated_properties": []string{"cartTotalMatches"}},
|
||||||
|
{"seed": 3, "exit_code": 0, "steps": 60, "duration_millis": 300000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 4, "exit_code": 0, "steps": 60, "duration_millis": 300000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 5, "exit_code": 0, "steps": 44, "duration_millis": 240000,
|
||||||
|
"first_violation_origin_step": 41, "violated_properties": []string{"cartTotalMatches", "backLeavesApp"}},
|
||||||
|
{"seed": 6, "exit_code": 0, "steps": 60, "duration_millis": 300000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 7, "exit_code": -1, "timed_out": true, "duration_millis": 900000},
|
||||||
|
{"seed": 8, "exit_code": 0, "steps": 60, "duration_millis": 300000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 9, "exit_code": 0, "steps": 60, "duration_millis": 300000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 10, "exit_code": 0, "steps": 21, "duration_millis": 130000,
|
||||||
|
"first_violation_origin_step": 19, "violated_properties": []string{"cartTotalMatches"}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func llmArmRecords() []map[string]any {
|
||||||
|
// Eight runs: six violate, one clean, one failed to launch.
|
||||||
|
return []map[string]any{
|
||||||
|
{"seed": 1, "exit_code": 0, "steps": 7, "duration_millis": 400000,
|
||||||
|
"first_violation_origin_step": 5, "violated_properties": []string{"cartTotalMatches"}},
|
||||||
|
{"seed": 2, "exit_code": 0, "steps": 9, "duration_millis": 420000,
|
||||||
|
"first_violation_origin_step": 8, "violated_properties": []string{"backLeavesApp"}},
|
||||||
|
{"seed": 3, "exit_code": 0, "steps": 60, "duration_millis": 1800000, "first_violation_origin_step": nil},
|
||||||
|
{"seed": 4, "exit_code": 0, "steps": 5, "duration_millis": 380000,
|
||||||
|
"first_violation_origin_step": 3, "violated_properties": []string{"cartTotalMatches"}},
|
||||||
|
{"seed": 5, "exit_code": 0, "steps": 13, "duration_millis": 500000,
|
||||||
|
"first_violation_origin_step": 11, "violated_properties": []string{"cartTotalMatches", "priceNeverNegative"}},
|
||||||
|
{"seed": 6, "exit_code": -1, "launch_error": "fork/exec sanderling: no such file or directory"},
|
||||||
|
{"seed": 7, "exit_code": 0, "steps": 6, "duration_millis": 390000,
|
||||||
|
"first_violation_origin_step": 6, "violated_properties": []string{"cartTotalMatches"}},
|
||||||
|
{"seed": 8, "exit_code": 0, "steps": 16, "duration_millis": 520000,
|
||||||
|
"first_violation_origin_step": 15, "violated_properties": []string{"backLeavesApp"}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_EndToEndOverFixtureCampaignDirectories(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
seededDirectory := filepath.Join(root, "seeded-web")
|
||||||
|
llmDirectory := filepath.Join(root, "llm-web")
|
||||||
|
buildFixtureCampaign(t, seededDirectory, "seeded", 60, seededArmRecords())
|
||||||
|
buildFixtureCampaign(t, llmDirectory, "llm", 60, llmArmRecords())
|
||||||
|
summaryPath := filepath.Join(root, "analysis.json")
|
||||||
|
|
||||||
|
var stdout bytes.Buffer
|
||||||
|
if err := run([]string{"--json", summaryPath, seededDirectory, llmDirectory}, &stdout, io.Discard); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
text := stdout.String()
|
||||||
|
for _, fragment := range []string{
|
||||||
|
"steps to first violation, right-censored at the step budget",
|
||||||
|
"log-rank across 2 arms",
|
||||||
|
"pairwise wilcoxon rank-sum",
|
||||||
|
"llm vs seeded",
|
||||||
|
"excluded 1 run(s) as missing data",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(text, fragment) {
|
||||||
|
t.Errorf("stdout is missing %q\n%s", fragment, text)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := os.ReadFile(summaryPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var result analysis
|
||||||
|
if err := json.Unmarshal(body, &result); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(result.Arms) != 2 {
|
||||||
|
t.Fatalf("%d arms, want 2", len(result.Arms))
|
||||||
|
}
|
||||||
|
byName := map[string]armSummary{}
|
||||||
|
for _, summary := range result.Arms {
|
||||||
|
byName[summary.Arm] = summary
|
||||||
|
}
|
||||||
|
|
||||||
|
seeded := byName["seeded"]
|
||||||
|
if seeded.Usable != 9 || seeded.Violated != 3 || seeded.Censored != 6 || seeded.Excluded != 1 {
|
||||||
|
t.Errorf("seeded arm %+v, want 9 usable, 3 violated, 6 censored, 1 excluded", seeded)
|
||||||
|
}
|
||||||
|
if seeded.ExcludedByReason[reasonTimedOut] != 1 {
|
||||||
|
t.Errorf("seeded exclusions %v, want one timeout", seeded.ExcludedByReason)
|
||||||
|
}
|
||||||
|
if seeded.MedianStepsToFirstViolation != nil {
|
||||||
|
t.Errorf("seeded median %v, want undefined with 3 of 9 violating",
|
||||||
|
*seeded.MedianStepsToFirstViolation)
|
||||||
|
}
|
||||||
|
if math.Abs(*seeded.ViolationRate-3.0/9.0) > 1e-12 {
|
||||||
|
t.Errorf("seeded violation rate %v, want 1/3", *seeded.ViolationRate)
|
||||||
|
}
|
||||||
|
// cartTotalMatches in 3 runs, backLeavesApp in 1 of 2 distinct defects.
|
||||||
|
if seeded.DistinctDefects != 2 || seeded.SingletonDefects != 1 {
|
||||||
|
t.Errorf("seeded defects %d singletons %d, want 2 and 1", seeded.DistinctDefects, seeded.SingletonDefects)
|
||||||
|
}
|
||||||
|
if seeded.TotalActions != 443 {
|
||||||
|
t.Errorf("seeded actions %d, want 443", seeded.TotalActions)
|
||||||
|
}
|
||||||
|
|
||||||
|
llm := byName["llm"]
|
||||||
|
if llm.Usable != 7 || llm.Violated != 6 || llm.Censored != 1 || llm.Excluded != 1 {
|
||||||
|
t.Errorf("llm arm %+v, want 7 usable, 6 violated, 1 censored, 1 excluded", llm)
|
||||||
|
}
|
||||||
|
if llm.ExcludedByReason[reasonLaunchError] != 1 {
|
||||||
|
t.Errorf("llm exclusions %v, want one launch error", llm.ExcludedByReason)
|
||||||
|
}
|
||||||
|
if llm.MedianStepsToFirstViolation == nil || *llm.MedianStepsToFirstViolation != 8 {
|
||||||
|
t.Errorf("llm median %v, want 8", llm.MedianStepsToFirstViolation)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.LogRank == nil {
|
||||||
|
t.Fatal("no log-rank result")
|
||||||
|
}
|
||||||
|
if result.LogRank.DegreesOfFreedom != 1 {
|
||||||
|
t.Errorf("log-rank df %d, want 1", result.LogRank.DegreesOfFreedom)
|
||||||
|
}
|
||||||
|
if result.LogRank.PValue > 0.05 {
|
||||||
|
t.Errorf("log-rank p %v, want the two clearly different arms to separate", result.LogRank.PValue)
|
||||||
|
}
|
||||||
|
if len(result.Pairwise) != 1 {
|
||||||
|
t.Fatalf("%d comparisons, want 1", len(result.Pairwise))
|
||||||
|
}
|
||||||
|
pair := result.Pairwise[0]
|
||||||
|
if pair.First != "llm" || pair.Second != "seeded" {
|
||||||
|
t.Errorf("comparison %s vs %s, want arms in sorted order", pair.First, pair.Second)
|
||||||
|
}
|
||||||
|
if pair.A12 >= 0.5 {
|
||||||
|
t.Errorf("a12 %v, want the arm that violates sooner below 0.5", pair.A12)
|
||||||
|
}
|
||||||
|
if pair.HolmPValue != pair.PValue {
|
||||||
|
t.Errorf("holm p %v differs from raw p %v in a family of one", pair.HolmPValue, pair.PValue)
|
||||||
|
}
|
||||||
|
if pair.Exact {
|
||||||
|
t.Error("used the exact null distribution despite the tie mass at the budget")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_ReportsBothArmsWhenOneHasNothingUsable(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
good := filepath.Join(root, "good")
|
||||||
|
broken := filepath.Join(root, "broken")
|
||||||
|
buildFixtureCampaign(t, good, "good", 30, []map[string]any{
|
||||||
|
{"seed": 1, "exit_code": 0, "steps": 30},
|
||||||
|
{"seed": 2, "exit_code": 0, "steps": 9, "duration_millis": 1000,
|
||||||
|
"first_violation_origin_step": 9, "violated_properties": []string{"cartTotalMatches"}},
|
||||||
|
})
|
||||||
|
buildFixtureCampaign(t, broken, "broken", 30, []map[string]any{
|
||||||
|
{"seed": 1, "exit_code": 3},
|
||||||
|
{"seed": 2, "timed_out": true, "exit_code": -1},
|
||||||
|
})
|
||||||
|
|
||||||
|
var stdout bytes.Buffer
|
||||||
|
if err := run([]string{good, broken}, &stdout, io.Discard); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
text := stdout.String()
|
||||||
|
if !strings.Contains(text, "broken") {
|
||||||
|
t.Errorf("the arm with nothing usable is not reported\n%s", text)
|
||||||
|
}
|
||||||
|
if strings.Contains(text, "log-rank across") {
|
||||||
|
t.Errorf("ran a log-rank with only one testable arm\n%s", text)
|
||||||
|
}
|
||||||
|
if !strings.Contains(text, "arms with no usable runs are reported but left out") {
|
||||||
|
t.Errorf("no note about the dropped arm\n%s", text)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_JsonToStdout(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
directory := filepath.Join(root, "only")
|
||||||
|
buildFixtureCampaign(t, directory, "only", 20, []map[string]any{
|
||||||
|
{"seed": 1, "exit_code": 0, "steps": 20},
|
||||||
|
})
|
||||||
|
var stdout bytes.Buffer
|
||||||
|
if err := run([]string{"--json", "-", "--campaign", directory}, &stdout, io.Discard); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
start := strings.Index(stdout.String(), "{")
|
||||||
|
if start < 0 {
|
||||||
|
t.Fatalf("no json in stdout\n%s", stdout.String())
|
||||||
|
}
|
||||||
|
var result analysis
|
||||||
|
if err := json.Unmarshal([]byte(stdout.String()[start:]), &result); err != nil {
|
||||||
|
t.Fatalf("json: %v", err)
|
||||||
|
}
|
||||||
|
if len(result.Arms) != 1 || result.Arms[0].Arm != "only" {
|
||||||
|
t.Errorf("arms %+v", result.Arms)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_RejectsTheSameDirectoryTwice(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
directory := filepath.Join(root, "one")
|
||||||
|
buildFixtureCampaign(t, directory, "one", 20, []map[string]any{{"seed": 1, "exit_code": 0, "steps": 20}})
|
||||||
|
err := run([]string{directory, directory}, io.Discard, io.Discard)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "twice") {
|
||||||
|
t.Fatalf("error %v, want a refusal to double count", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_RequiresACampaignDirectory(t *testing.T) {
|
||||||
|
if err := run(nil, io.Discard, io.Discard); err == nil {
|
||||||
|
t.Fatal("expected an error with no campaign directories")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
// Published right-censored datasets whose log-rank and Kaplan-Meier results are
|
||||||
|
// reported in the survival-analysis literature and in R's survival package, so
|
||||||
|
// every expected number in these tests can be checked against a source rather
|
||||||
|
// than against this tool's own output.
|
||||||
|
|
||||||
|
// gehanSixMercaptopurine and gehanPlacebo are remission times in weeks from the
|
||||||
|
// 6-MP versus placebo trial in acute leukaemia, Freireich et al. (1963). This is
|
||||||
|
// the dataset R's survival literature calls gehan. A trailing plus in the
|
||||||
|
// published listing marks a censored time.
|
||||||
|
//
|
||||||
|
// 6-MP: 6, 6, 6, 6+, 7, 9+, 10, 10+, 11+, 13, 16, 17+, 19+, 20+, 22, 23, 25+, 32+, 32+, 34+, 35+
|
||||||
|
// placebo: 1, 1, 2, 2, 3, 4, 4, 5, 5, 8, 8, 8, 8, 11, 11, 12, 12, 15, 17, 22, 23
|
||||||
|
var (
|
||||||
|
gehanSixMercaptopurine = []observation{
|
||||||
|
{6, true}, {6, true}, {6, true}, {6, false},
|
||||||
|
{7, true}, {9, false}, {10, true}, {10, false},
|
||||||
|
{11, false}, {13, true}, {16, true}, {17, false},
|
||||||
|
{19, false}, {20, false}, {22, true}, {23, true},
|
||||||
|
{25, false}, {32, false}, {32, false}, {34, false}, {35, false},
|
||||||
|
}
|
||||||
|
gehanPlacebo = []observation{
|
||||||
|
{1, true}, {1, true}, {2, true}, {2, true}, {3, true},
|
||||||
|
{4, true}, {4, true}, {5, true}, {5, true}, {8, true},
|
||||||
|
{8, true}, {8, true}, {8, true}, {11, true}, {11, true},
|
||||||
|
{12, true}, {12, true}, {15, true}, {17, true}, {22, true}, {23, true},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
// amlMaintained and amlNonmaintained are the acute myelogenous leukaemia
|
||||||
|
// survival times in weeks from Miller (1997), shipped as the aml dataset in R's
|
||||||
|
// survival package. Five subjects are censored, at 13, 16, 28, 45 and 161 weeks.
|
||||||
|
//
|
||||||
|
// maintained: 9, 13, 13+, 18, 23, 28+, 31, 34, 45+, 48, 161+
|
||||||
|
// nonmaintained: 5, 5, 8, 8, 12, 16+, 23, 27, 30, 33, 43, 45
|
||||||
|
var (
|
||||||
|
amlMaintained = []observation{
|
||||||
|
{9, true}, {13, true}, {13, false}, {18, true}, {23, true},
|
||||||
|
{28, false}, {31, true}, {34, true}, {45, false}, {48, true}, {161, false},
|
||||||
|
}
|
||||||
|
amlNonmaintained = []observation{
|
||||||
|
{5, true}, {5, true}, {8, true}, {8, true}, {12, true}, {16, false},
|
||||||
|
{23, true}, {27, true}, {30, true}, {33, true}, {43, true}, {45, true},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
// chorioamnionTerm and chorioamnionEarly are permeability constants of the human
|
||||||
|
// chorioamnion at term and between 12 and 26 weeks gestational age, Hollander
|
||||||
|
// and Wolfe (1973), 69f. R's wilcox.test help page uses exactly these vectors as
|
||||||
|
// its two-sample example.
|
||||||
|
var (
|
||||||
|
chorioamnionTerm = []float64{0.80, 0.83, 1.89, 1.04, 1.45, 1.38, 1.91, 1.64, 0.73, 1.46}
|
||||||
|
chorioamnionEarly = []float64{1.15, 0.88, 0.90, 0.74, 1.21}
|
||||||
|
)
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
|
)
|
||||||
|
|
||||||
|
// holm applies the Holm (1979) step-down correction within one family of
|
||||||
|
// comparisons, enforcing monotonicity across the sorted p-values the way R's
|
||||||
|
// p.adjust does. Holm, "A Simple Sequentially Rejective Multiple Test
|
||||||
|
// Procedure", Scandinavian Journal of Statistics 6(2), 65-70.
|
||||||
|
func holm(pValues []float64) []float64 {
|
||||||
|
count := len(pValues)
|
||||||
|
adjusted := make([]float64, count)
|
||||||
|
order := make([]int, count)
|
||||||
|
for index := range order {
|
||||||
|
order[index] = index
|
||||||
|
}
|
||||||
|
slices.SortStableFunc(order, func(left, right int) int {
|
||||||
|
switch {
|
||||||
|
case pValues[left] < pValues[right]:
|
||||||
|
return -1
|
||||||
|
case pValues[left] > pValues[right]:
|
||||||
|
return 1
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
})
|
||||||
|
running := 0.0
|
||||||
|
for position, index := range order {
|
||||||
|
scaled := float64(count-position) * pValues[index]
|
||||||
|
running = math.Max(running, scaled)
|
||||||
|
adjusted[index] = math.Min(running, 1)
|
||||||
|
}
|
||||||
|
return adjusted
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Both cases are printed R output in Eve Slavich, "Four strategies for dealing
|
||||||
|
// with multiple comparisons", UNSW Stats Central, slides 9 and 10:
|
||||||
|
//
|
||||||
|
// pValues = c(0.01, 0.2, 0.08, 0.03)
|
||||||
|
// p.adjust(pValues, method = "holm")
|
||||||
|
// ## [1] 0.04 0.20 0.16 0.09
|
||||||
|
//
|
||||||
|
// pValues = c(0.01, 0.2, 0.08, 0.03, 0.02, 0.01)
|
||||||
|
// p.adjust(pValues, method = "holm")
|
||||||
|
// ## [1] 0.06 0.20 0.16 0.09 0.08 0.06
|
||||||
|
//
|
||||||
|
// The second case exercises the monotonicity step: sorted p are
|
||||||
|
// .01 .01 .02 .03 .08 .20, scaled by 6 5 4 3 2 1 to .06 .05 .08 .09 .16 .20,
|
||||||
|
// and the running maximum lifts the second back to .06.
|
||||||
|
func TestHolm_MatchesPublishedAdjustment(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
raw []float64
|
||||||
|
expected []float64
|
||||||
|
}{
|
||||||
|
{[]float64{0.01, 0.2, 0.08, 0.03}, []float64{0.04, 0.20, 0.16, 0.09}},
|
||||||
|
{[]float64{0.01, 0.2, 0.08, 0.03, 0.02, 0.01}, []float64{0.06, 0.20, 0.16, 0.09, 0.08, 0.06}},
|
||||||
|
}
|
||||||
|
for _, test := range cases {
|
||||||
|
adjusted := holm(test.raw)
|
||||||
|
for index, want := range test.expected {
|
||||||
|
if math.Abs(adjusted[index]-want) > 1e-12 {
|
||||||
|
t.Errorf("holm(%v)[%d] = %v, want %v", test.raw, index, adjusted[index], want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHolm_CapsAtOneAndKeepsOrder(t *testing.T) {
|
||||||
|
adjusted := holm([]float64{0.4, 0.5, 0.9})
|
||||||
|
for index, value := range adjusted {
|
||||||
|
if value != 1 {
|
||||||
|
t.Errorf("adjusted[%d] = %v, want 1", index, value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
single := holm([]float64{0.03})
|
||||||
|
if len(single) != 1 || single[0] != 0.03 {
|
||||||
|
t.Errorf("single comparison adjusted to %v, want 0.03 unchanged", single)
|
||||||
|
}
|
||||||
|
if got := holm(nil); len(got) != 0 {
|
||||||
|
t.Errorf("holm(nil) = %v, want empty", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,242 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
manifestFileName = "campaign.json"
|
||||||
|
recordsFileName = "runs.jsonl"
|
||||||
|
maxRecordBytes = 4 * 1024 * 1024
|
||||||
|
)
|
||||||
|
|
||||||
|
// manifest mirrors the fields analyze reads from campaign.json. The step budget
|
||||||
|
// lives here rather than in any run, because it is what clean runs are censored
|
||||||
|
// at and every run in an arm has to share it.
|
||||||
|
type manifest struct {
|
||||||
|
Arm string `json:"arm"`
|
||||||
|
Generator string `json:"generator"`
|
||||||
|
Platform string `json:"platform"`
|
||||||
|
MaxSteps int `json:"max_steps"`
|
||||||
|
Seeds []int64 `json:"seeds"`
|
||||||
|
Host string `json:"host"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// runRecord mirrors the fields analyze reads from one line of runs.jsonl.
|
||||||
|
type runRecord struct {
|
||||||
|
Seed int64 `json:"seed"`
|
||||||
|
ExitCode int `json:"exit_code"`
|
||||||
|
LaunchError string `json:"launch_error"`
|
||||||
|
TimedOut bool `json:"timed_out"`
|
||||||
|
DurationMillis int64 `json:"duration_millis"`
|
||||||
|
TraceError string `json:"trace_error"`
|
||||||
|
Steps int `json:"steps"`
|
||||||
|
FirstViolationOriginStep *int `json:"first_violation_origin_step"`
|
||||||
|
ViolatedProperties []string `json:"violated_properties"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exclusion reasons. A run that failed or timed out is missing data, not a
|
||||||
|
// censored observation: it never ran its budget, so treating it as a clean run
|
||||||
|
// that survived to the budget would bias the survival estimate downward.
|
||||||
|
const (
|
||||||
|
reasonLaunchError = "launch error"
|
||||||
|
reasonTimedOut = "timed out"
|
||||||
|
reasonNonzeroExit = "nonzero exit"
|
||||||
|
reasonTraceError = "unreadable trace"
|
||||||
|
reasonMalformedStep = "violation step outside the budget"
|
||||||
|
)
|
||||||
|
|
||||||
|
type classifiedRun struct {
|
||||||
|
Seed int64
|
||||||
|
Steps int
|
||||||
|
DurationMillis int64
|
||||||
|
OriginStep int
|
||||||
|
Violated bool
|
||||||
|
ClampedToBudget bool
|
||||||
|
ViolatedProperties []string
|
||||||
|
ExcludedBecause string
|
||||||
|
}
|
||||||
|
|
||||||
|
type arm struct {
|
||||||
|
Name string
|
||||||
|
Budget int
|
||||||
|
Generator string
|
||||||
|
Platform string
|
||||||
|
Directories []string
|
||||||
|
Runs []classifiedRun
|
||||||
|
MissingSeeds []int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadCampaign(directory string) (manifest, []runRecord, error) {
|
||||||
|
body, err := os.ReadFile(filepath.Join(directory, manifestFileName))
|
||||||
|
if err != nil {
|
||||||
|
return manifest{}, nil, fmt.Errorf("read %s: %w", manifestFileName, err)
|
||||||
|
}
|
||||||
|
var declared manifest
|
||||||
|
if err := json.Unmarshal(body, &declared); err != nil {
|
||||||
|
return manifest{}, nil, fmt.Errorf("parse %s in %s: %w", manifestFileName, directory, err)
|
||||||
|
}
|
||||||
|
if declared.Arm == "" {
|
||||||
|
return manifest{}, nil, fmt.Errorf("%s in %s has no arm", manifestFileName, directory)
|
||||||
|
}
|
||||||
|
if declared.MaxSteps <= 0 {
|
||||||
|
return manifest{}, nil, fmt.Errorf("%s in %s has max_steps %d: clean runs have nothing to be censored at",
|
||||||
|
manifestFileName, directory, declared.MaxSteps)
|
||||||
|
}
|
||||||
|
|
||||||
|
file, err := os.Open(filepath.Join(directory, recordsFileName))
|
||||||
|
if err != nil {
|
||||||
|
return manifest{}, nil, fmt.Errorf("read %s: %w", recordsFileName, err)
|
||||||
|
}
|
||||||
|
defer file.Close()
|
||||||
|
|
||||||
|
var records []runRecord
|
||||||
|
scanner := bufio.NewScanner(file)
|
||||||
|
scanner.Buffer(make([]byte, 0, 64*1024), maxRecordBytes)
|
||||||
|
lineNumber := 0
|
||||||
|
for scanner.Scan() {
|
||||||
|
lineNumber++
|
||||||
|
raw := strings.TrimSpace(scanner.Text())
|
||||||
|
if raw == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var record runRecord
|
||||||
|
if err := json.Unmarshal([]byte(raw), &record); err != nil {
|
||||||
|
return manifest{}, nil, fmt.Errorf("%s line %d in %s: %w", recordsFileName, lineNumber, directory, err)
|
||||||
|
}
|
||||||
|
records = append(records, record)
|
||||||
|
}
|
||||||
|
if err := scanner.Err(); err != nil {
|
||||||
|
return manifest{}, nil, fmt.Errorf("read %s in %s: %w", recordsFileName, directory, err)
|
||||||
|
}
|
||||||
|
return declared, records, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// classify turns one record into the run the analysis works with, deciding
|
||||||
|
// whether it is usable and, if it is, whether it is an event or censored.
|
||||||
|
func classify(record runRecord, budget int) classifiedRun {
|
||||||
|
item := classifiedRun{
|
||||||
|
Seed: record.Seed,
|
||||||
|
Steps: record.Steps,
|
||||||
|
DurationMillis: record.DurationMillis,
|
||||||
|
ViolatedProperties: slices.Clone(record.ViolatedProperties),
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case record.LaunchError != "":
|
||||||
|
item.ExcludedBecause = reasonLaunchError
|
||||||
|
return item
|
||||||
|
case record.TimedOut:
|
||||||
|
item.ExcludedBecause = reasonTimedOut
|
||||||
|
return item
|
||||||
|
case record.ExitCode != 0:
|
||||||
|
item.ExcludedBecause = reasonNonzeroExit
|
||||||
|
return item
|
||||||
|
case record.TraceError != "":
|
||||||
|
item.ExcludedBecause = reasonTraceError
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
if record.FirstViolationOriginStep == nil {
|
||||||
|
if len(record.ViolatedProperties) > 0 {
|
||||||
|
item.ExcludedBecause = reasonMalformedStep
|
||||||
|
}
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
origin := *record.FirstViolationOriginStep
|
||||||
|
if origin < 1 {
|
||||||
|
item.ExcludedBecause = reasonMalformedStep
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
item.Violated = true
|
||||||
|
if origin > budget {
|
||||||
|
// The run-end finalize line reports obligations that never discharged
|
||||||
|
// at an index one past the last executed step. That is a real detection
|
||||||
|
// but not a real step, so it is held at the budget and counted.
|
||||||
|
origin = budget
|
||||||
|
item.ClampedToBudget = true
|
||||||
|
}
|
||||||
|
item.OriginStep = origin
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
|
||||||
|
// groupArms folds every campaign directory into its arm. Two directories with
|
||||||
|
// the same arm label are pooled, which is how a campaign split across hosts is
|
||||||
|
// analysed, but they must agree on the step budget.
|
||||||
|
func groupArms(directories []string) ([]arm, error) {
|
||||||
|
byName := map[string]*arm{}
|
||||||
|
var order []string
|
||||||
|
for _, directory := range directories {
|
||||||
|
declared, records, err := loadCampaign(directory)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
current, seen := byName[declared.Arm]
|
||||||
|
if !seen {
|
||||||
|
current = &arm{
|
||||||
|
Name: declared.Arm,
|
||||||
|
Budget: declared.MaxSteps,
|
||||||
|
Generator: declared.Generator,
|
||||||
|
Platform: declared.Platform,
|
||||||
|
}
|
||||||
|
byName[declared.Arm] = current
|
||||||
|
order = append(order, declared.Arm)
|
||||||
|
}
|
||||||
|
if current.Budget != declared.MaxSteps {
|
||||||
|
return nil, fmt.Errorf("arm %q has step budget %d in an earlier campaign and %d in %s: "+
|
||||||
|
"runs censored at different budgets cannot be pooled",
|
||||||
|
declared.Arm, current.Budget, declared.MaxSteps, directory)
|
||||||
|
}
|
||||||
|
current.Directories = append(current.Directories, directory)
|
||||||
|
|
||||||
|
present := map[int64]bool{}
|
||||||
|
for _, record := range records {
|
||||||
|
present[record.Seed] = true
|
||||||
|
current.Runs = append(current.Runs, classify(record, declared.MaxSteps))
|
||||||
|
}
|
||||||
|
for _, seed := range declared.Seeds {
|
||||||
|
if !present[seed] {
|
||||||
|
current.MissingSeeds = append(current.MissingSeeds, seed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
slices.Sort(order)
|
||||||
|
arms := make([]arm, 0, len(order))
|
||||||
|
for _, name := range order {
|
||||||
|
arms = append(arms, *byName[name])
|
||||||
|
}
|
||||||
|
return arms, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// observations returns the usable runs as survival data: an event at the step
|
||||||
|
// that armed the first violation, or a censored observation at the step budget.
|
||||||
|
func (a arm) observations() []observation {
|
||||||
|
var result []observation
|
||||||
|
for _, item := range a.Runs {
|
||||||
|
if item.ExcludedBecause != "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if item.Violated {
|
||||||
|
result = append(result, observation{Steps: float64(item.OriginStep), Event: true})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
result = append(result, observation{Steps: float64(a.Budget), Event: false})
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// stepTimes is the observations flattened to plain numbers, censored runs held
|
||||||
|
// at the budget. Holding them there rather than dropping them is conservative:
|
||||||
|
// it can only understate how much sooner a violating arm finds its first
|
||||||
|
// defect, never overstate it.
|
||||||
|
func (a arm) stepTimes() []float64 {
|
||||||
|
var result []float64
|
||||||
|
for _, item := range a.observations() {
|
||||||
|
result = append(result, item.Steps)
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func stepPointer(value int) *int { return &value }
|
||||||
|
|
||||||
|
func TestClassify_FailedAndTimedOutRunsAreMissingDataNotCensored(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
record runRecord
|
||||||
|
reason string
|
||||||
|
}{
|
||||||
|
{"launch error", runRecord{LaunchError: "fork/exec: no such file"}, reasonLaunchError},
|
||||||
|
{"timed out", runRecord{TimedOut: true, ExitCode: -1}, reasonTimedOut},
|
||||||
|
{"nonzero exit", runRecord{ExitCode: 3}, reasonNonzeroExit},
|
||||||
|
{"unreadable trace", runRecord{TraceError: "no run directory with meta.json"}, reasonTraceError},
|
||||||
|
{"violation at step zero", runRecord{FirstViolationOriginStep: stepPointer(0)}, reasonMalformedStep},
|
||||||
|
{"violation without a step", runRecord{ViolatedProperties: []string{"cartTotal"}}, reasonMalformedStep},
|
||||||
|
}
|
||||||
|
for _, test := range cases {
|
||||||
|
item := classify(test.record, 50)
|
||||||
|
if item.ExcludedBecause != test.reason {
|
||||||
|
t.Errorf("%s: excluded because %q, want %q", test.name, item.ExcludedBecause, test.reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassify_CleanRunIsCensoredAtTheBudget(t *testing.T) {
|
||||||
|
item := classify(runRecord{Seed: 4, Steps: 50, DurationMillis: 1000}, 50)
|
||||||
|
if item.ExcludedBecause != "" {
|
||||||
|
t.Fatalf("excluded because %q", item.ExcludedBecause)
|
||||||
|
}
|
||||||
|
if item.Violated {
|
||||||
|
t.Error("clean run marked as violated")
|
||||||
|
}
|
||||||
|
current := arm{Budget: 50, Runs: []classifiedRun{item}}
|
||||||
|
observations := current.observations()
|
||||||
|
if len(observations) != 1 || observations[0].Event || observations[0].Steps != 50 {
|
||||||
|
t.Errorf("observations %+v, want one censored observation at 50", observations)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClassify_ViolationIsAnEventAtTheOriginStep(t *testing.T) {
|
||||||
|
item := classify(runRecord{Seed: 5, Steps: 12, FirstViolationOriginStep: stepPointer(7)}, 50)
|
||||||
|
if !item.Violated || item.OriginStep != 7 || item.ClampedToBudget {
|
||||||
|
t.Fatalf("run %+v, want an unclamped event at step 7", item)
|
||||||
|
}
|
||||||
|
current := arm{Budget: 50, Runs: []classifiedRun{item}}
|
||||||
|
observations := current.observations()
|
||||||
|
if len(observations) != 1 || !observations[0].Event || observations[0].Steps != 7 {
|
||||||
|
t.Errorf("observations %+v, want one event at 7", observations)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The run-end finalize line reports at an index one past the last executed step,
|
||||||
|
// so an origin past the budget is held at the budget and counted rather than
|
||||||
|
// silently turned into a censored run.
|
||||||
|
func TestClassify_ViolationPastTheBudgetIsHeldAtTheBudget(t *testing.T) {
|
||||||
|
item := classify(runRecord{FirstViolationOriginStep: stepPointer(51)}, 50)
|
||||||
|
if !item.Violated || item.OriginStep != 50 || !item.ClampedToBudget {
|
||||||
|
t.Errorf("run %+v, want a clamped event at 50", item)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeCampaign(t *testing.T, directory string, declared map[string]any, records []map[string]any) {
|
||||||
|
t.Helper()
|
||||||
|
if err := os.MkdirAll(directory, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
body, err := json.MarshalIndent(declared, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(filepath.Join(directory, manifestFileName), append(body, '\n'), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var lines strings.Builder
|
||||||
|
for _, record := range records {
|
||||||
|
line, err := json.Marshal(record)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
lines.Write(line)
|
||||||
|
lines.WriteByte('\n')
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(filepath.Join(directory, recordsFileName), []byte(lines.String()), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGroupArms_PoolsDirectoriesSharingAnArmAndReportsMissingSeeds(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
writeCampaign(t, filepath.Join(root, "north"), map[string]any{
|
||||||
|
"arm": "seeded", "max_steps": 40, "seeds": []int{1, 2, 3},
|
||||||
|
}, []map[string]any{
|
||||||
|
{"seed": 1, "exit_code": 0, "steps": 40},
|
||||||
|
{"seed": 2, "exit_code": 0, "steps": 9, "first_violation_origin_step": 9},
|
||||||
|
})
|
||||||
|
writeCampaign(t, filepath.Join(root, "south"), map[string]any{
|
||||||
|
"arm": "seeded", "max_steps": 40, "seeds": []int{4},
|
||||||
|
}, []map[string]any{
|
||||||
|
{"seed": 4, "exit_code": 0, "steps": 40},
|
||||||
|
})
|
||||||
|
|
||||||
|
arms, err := groupArms([]string{filepath.Join(root, "north"), filepath.Join(root, "south")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(arms) != 1 {
|
||||||
|
t.Fatalf("%d arms, want 1", len(arms))
|
||||||
|
}
|
||||||
|
if len(arms[0].Runs) != 3 {
|
||||||
|
t.Errorf("%d runs, want 3", len(arms[0].Runs))
|
||||||
|
}
|
||||||
|
if len(arms[0].MissingSeeds) != 1 || arms[0].MissingSeeds[0] != 3 {
|
||||||
|
t.Errorf("missing seeds %v, want [3]", arms[0].MissingSeeds)
|
||||||
|
}
|
||||||
|
if len(arms[0].Directories) != 2 {
|
||||||
|
t.Errorf("directories %v, want both", arms[0].Directories)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGroupArms_RejectsDisagreeingStepBudgets(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
writeCampaign(t, filepath.Join(root, "a"), map[string]any{"arm": "seeded", "max_steps": 40, "seeds": []int{1}},
|
||||||
|
[]map[string]any{{"seed": 1, "exit_code": 0, "steps": 40}})
|
||||||
|
writeCampaign(t, filepath.Join(root, "b"), map[string]any{"arm": "seeded", "max_steps": 80, "seeds": []int{2}},
|
||||||
|
[]map[string]any{{"seed": 2, "exit_code": 0, "steps": 80}})
|
||||||
|
|
||||||
|
_, err := groupArms([]string{filepath.Join(root, "a"), filepath.Join(root, "b")})
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "different budgets") {
|
||||||
|
t.Fatalf("error %v, want a refusal to pool different budgets", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGroupArms_RejectsAMissingStepBudget(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
writeCampaign(t, filepath.Join(root, "a"), map[string]any{"arm": "seeded", "seeds": []int{1}},
|
||||||
|
[]map[string]any{{"seed": 1, "exit_code": 0}})
|
||||||
|
_, err := groupArms([]string{filepath.Join(root, "a")})
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "censored at") {
|
||||||
|
t.Fatalf("error %v, want a complaint about max_steps", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGroupArms_ReportsBadRecordLines(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
directory := filepath.Join(root, "a")
|
||||||
|
writeCampaign(t, directory, map[string]any{"arm": "seeded", "max_steps": 40, "seeds": []int{1}}, nil)
|
||||||
|
if err := os.WriteFile(filepath.Join(directory, recordsFileName), []byte("{\"seed\":1}\nnot json\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, err := groupArms([]string{directory})
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "line 2") {
|
||||||
|
t.Fatalf("error %v, want the offending line number", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,102 @@
|
|||||||
|
// Command analyze reduces campaign directories to the statistics the
|
||||||
|
// evaluation reports. The primary outcome is steps to first violation, with
|
||||||
|
// clean runs right-censored at the step budget rather than discarded: defect
|
||||||
|
// yield per run is a binary that would need on the order of eighty runs an arm
|
||||||
|
// to separate, while survival analysis uses every run, including the clean ones.
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"flag"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const usage = `analyze reports the statistics of a sanderling evaluation from campaign directories.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
analyze [--json <path>] <campaign-dir> [<campaign-dir> ...]
|
||||||
|
|
||||||
|
Each directory is one produced by the campaign tool and must hold campaign.json
|
||||||
|
and runs.jsonl. Directories sharing an arm label are pooled and must agree on
|
||||||
|
the step budget.
|
||||||
|
`
|
||||||
|
|
||||||
|
type stringList []string
|
||||||
|
|
||||||
|
func (list *stringList) String() string { return strings.Join(*list, ",") }
|
||||||
|
|
||||||
|
func (list *stringList) Set(value string) error {
|
||||||
|
if strings.TrimSpace(value) == "" {
|
||||||
|
return errors.New("empty campaign directory")
|
||||||
|
}
|
||||||
|
*list = append(*list, value)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func run(arguments []string, stdout, stderr io.Writer) error {
|
||||||
|
flagSet := flag.NewFlagSet("analyze", flag.ContinueOnError)
|
||||||
|
flagSet.SetOutput(stderr)
|
||||||
|
flagSet.Usage = func() {
|
||||||
|
fmt.Fprint(stderr, usage)
|
||||||
|
flagSet.PrintDefaults()
|
||||||
|
}
|
||||||
|
var directories stringList
|
||||||
|
var jsonPath string
|
||||||
|
flagSet.Var(&directories, "campaign", "campaign directory to read; repeat for more, or pass them as arguments")
|
||||||
|
flagSet.StringVar(&jsonPath, "json", "", "write the machine-readable summary here, or - for stdout")
|
||||||
|
if err := flagSet.Parse(arguments); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
directories = append(directories, flagSet.Args()...)
|
||||||
|
if len(directories) == 0 {
|
||||||
|
return errors.New("no campaign directories given")
|
||||||
|
}
|
||||||
|
seen := map[string]bool{}
|
||||||
|
for _, directory := range directories {
|
||||||
|
resolved, err := filepath.Abs(directory)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("resolve %s: %w", directory, err)
|
||||||
|
}
|
||||||
|
if seen[resolved] {
|
||||||
|
return fmt.Errorf("campaign directory %s given twice: its runs would be counted twice", directory)
|
||||||
|
}
|
||||||
|
seen[resolved] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
arms, err := groupArms(directories)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
result := analyse(arms, time.Now().UTC())
|
||||||
|
writeReport(result, stdout)
|
||||||
|
|
||||||
|
if jsonPath == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
body, err := json.MarshalIndent(result, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("marshal summary: %w", err)
|
||||||
|
}
|
||||||
|
body = append(body, '\n')
|
||||||
|
if jsonPath == "-" {
|
||||||
|
_, err = stdout.Write(body)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return os.WriteFile(jsonPath, body, 0o644)
|
||||||
|
}
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
if err := run(os.Args[1:], os.Stdout, os.Stderr); err != nil {
|
||||||
|
if errors.Is(err, flag.ErrHelp) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fmt.Fprintf(os.Stderr, "error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,220 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
|
)
|
||||||
|
|
||||||
|
type rankSumResult struct {
|
||||||
|
FirstSize int `json:"first_size"`
|
||||||
|
SecondSize int `json:"second_size"`
|
||||||
|
Statistic float64 `json:"mann_whitney_u"`
|
||||||
|
A12 float64 `json:"a12"`
|
||||||
|
PValue float64 `json:"p_value"`
|
||||||
|
Exact bool `json:"exact"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// exactRankSumLimit matches R's wilcox.test: the exact null distribution is
|
||||||
|
// used only when both samples are below this size and nothing is tied.
|
||||||
|
const exactRankSumLimit = 50
|
||||||
|
|
||||||
|
// vargaDelaneyA12 is the probability that a value drawn from first exceeds one
|
||||||
|
// drawn from second, counting a tie as half:
|
||||||
|
//
|
||||||
|
// A = P(X > Y) + 0.5 * P(X = Y)
|
||||||
|
//
|
||||||
|
// Vargha and Delaney (2000), "A Critique and Improvement of the CL Common
|
||||||
|
// Language Effect Size Statistics of McGraw and Wong", Journal of Educational
|
||||||
|
// and Behavioral Statistics 25(2), 101-132.
|
||||||
|
func vargaDelaneyA12(first, second []float64) float64 {
|
||||||
|
if len(first) == 0 || len(second) == 0 {
|
||||||
|
return math.NaN()
|
||||||
|
}
|
||||||
|
total := 0.0
|
||||||
|
for _, left := range first {
|
||||||
|
for _, right := range second {
|
||||||
|
switch {
|
||||||
|
case left > right:
|
||||||
|
total++
|
||||||
|
case left == right:
|
||||||
|
total += 0.5
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return total / float64(len(first)*len(second))
|
||||||
|
}
|
||||||
|
|
||||||
|
// rankSum is the two-sided Wilcoxon rank-sum (Mann-Whitney) test. The reported
|
||||||
|
// statistic is U for the first sample, the same quantity R's wilcox.test calls
|
||||||
|
// W. The exact null distribution is used when there are no ties and both
|
||||||
|
// samples are small; otherwise the normal approximation is used with the
|
||||||
|
// continuity correction and the tie correction to the variance.
|
||||||
|
func rankSum(first, second []float64) rankSumResult {
|
||||||
|
firstSize, secondSize := len(first), len(second)
|
||||||
|
result := rankSumResult{
|
||||||
|
FirstSize: firstSize,
|
||||||
|
SecondSize: secondSize,
|
||||||
|
Statistic: math.NaN(),
|
||||||
|
A12: math.NaN(),
|
||||||
|
PValue: math.NaN(),
|
||||||
|
}
|
||||||
|
if firstSize == 0 || secondSize == 0 {
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
pooled := make([]float64, 0, firstSize+secondSize)
|
||||||
|
pooled = append(pooled, first...)
|
||||||
|
pooled = append(pooled, second...)
|
||||||
|
ranks, tieGroups := midRanks(pooled)
|
||||||
|
|
||||||
|
rankTotal := 0.0
|
||||||
|
for index := 0; index < firstSize; index++ {
|
||||||
|
rankTotal += ranks[index]
|
||||||
|
}
|
||||||
|
statistic := rankTotal - float64(firstSize)*float64(firstSize+1)/2
|
||||||
|
result.Statistic = statistic
|
||||||
|
result.A12 = vargaDelaneyA12(first, second)
|
||||||
|
|
||||||
|
if len(tieGroups) == 0 && firstSize < exactRankSumLimit && secondSize < exactRankSumLimit {
|
||||||
|
result.Exact = true
|
||||||
|
result.PValue = exactRankSumTwoSided(statistic, firstSize, secondSize)
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
result.PValue = normalRankSumTwoSided(statistic, firstSize, secondSize, tieGroups)
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// normalRankSumTwoSided follows the large-sample branch of R's wilcox.test:
|
||||||
|
//
|
||||||
|
// sigma^2 = (m*n/12) * ((N+1) - sum(t^3 - t) / (N*(N-1)))
|
||||||
|
//
|
||||||
|
// where t runs over the sizes of the tied groups. The 0.5 shift toward the null
|
||||||
|
// mean is the continuity correction.
|
||||||
|
func normalRankSumTwoSided(statistic float64, firstSize, secondSize int, tieGroups []int) float64 {
|
||||||
|
sizeProduct := float64(firstSize) * float64(secondSize)
|
||||||
|
variance := rankSumVariance(firstSize, secondSize, tieGroups)
|
||||||
|
if variance <= 0 {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
centered := statistic - sizeProduct/2
|
||||||
|
correction := 0.0
|
||||||
|
switch {
|
||||||
|
case centered > 0:
|
||||||
|
correction = 0.5
|
||||||
|
case centered < 0:
|
||||||
|
correction = -0.5
|
||||||
|
}
|
||||||
|
z := (centered - correction) / math.Sqrt(variance)
|
||||||
|
tail := math.Min(standardNormalUpperTail(z), standardNormalUpperTail(-z))
|
||||||
|
return math.Min(2*tail, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func rankSumVariance(firstSize, secondSize int, tieGroups []int) float64 {
|
||||||
|
sizeProduct := float64(firstSize) * float64(secondSize)
|
||||||
|
total := float64(firstSize + secondSize)
|
||||||
|
tieAdjustment := 0.0
|
||||||
|
for _, size := range tieGroups {
|
||||||
|
count := float64(size)
|
||||||
|
tieAdjustment += count*count*count - count
|
||||||
|
}
|
||||||
|
return (sizeProduct / 12) * ((total + 1) - tieAdjustment/(total*(total-1)))
|
||||||
|
}
|
||||||
|
|
||||||
|
// exactRankSumTwoSided doubles the smaller exact tail, as R's wilcox.test does.
|
||||||
|
func exactRankSumTwoSided(statistic float64, firstSize, secondSize int) float64 {
|
||||||
|
counts := exactRankSumCounts(firstSize, secondSize)
|
||||||
|
total := 0.0
|
||||||
|
for _, count := range counts {
|
||||||
|
total += count
|
||||||
|
}
|
||||||
|
value := int(math.Round(statistic))
|
||||||
|
tail := 0.0
|
||||||
|
if statistic > float64(firstSize*secondSize)/2 {
|
||||||
|
for u := value; u < len(counts); u++ {
|
||||||
|
tail += counts[u]
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
for u := 0; u <= value && u < len(counts); u++ {
|
||||||
|
tail += counts[u]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return math.Min(2*tail/total, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// exactRankSumUpperTail is P(U >= statistic) under the null with no ties.
|
||||||
|
func exactRankSumUpperTail(statistic float64, firstSize, secondSize int) float64 {
|
||||||
|
counts := exactRankSumCounts(firstSize, secondSize)
|
||||||
|
total, tail := 0.0, 0.0
|
||||||
|
for u, count := range counts {
|
||||||
|
total += count
|
||||||
|
if float64(u) >= statistic {
|
||||||
|
tail += count
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return tail / total
|
||||||
|
}
|
||||||
|
|
||||||
|
// exactRankSumCounts returns the number of untied assignments producing each
|
||||||
|
// value of U from 0 to firstSize*secondSize. U equals the sum of the zero-based
|
||||||
|
// pooled positions held by the first sample, less firstSize*(firstSize-1)/2, so
|
||||||
|
// the count is a subset-sum tally over those positions.
|
||||||
|
func exactRankSumCounts(firstSize, secondSize int) []float64 {
|
||||||
|
maximum := firstSize * secondSize
|
||||||
|
offset := firstSize * (firstSize - 1) / 2
|
||||||
|
high := maximum + offset
|
||||||
|
table := make([][]float64, firstSize+1)
|
||||||
|
for index := range table {
|
||||||
|
table[index] = make([]float64, high+1)
|
||||||
|
}
|
||||||
|
table[0][0] = 1
|
||||||
|
for position := 0; position < firstSize+secondSize; position++ {
|
||||||
|
for chosen := min(position+1, firstSize); chosen >= 1; chosen-- {
|
||||||
|
row, previous := table[chosen], table[chosen-1]
|
||||||
|
for sum := high; sum >= position; sum-- {
|
||||||
|
if previous[sum-position] != 0 {
|
||||||
|
row[sum] += previous[sum-position]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
counts := make([]float64, maximum+1)
|
||||||
|
for u := range counts {
|
||||||
|
counts[u] = table[firstSize][u+offset]
|
||||||
|
}
|
||||||
|
return counts
|
||||||
|
}
|
||||||
|
|
||||||
|
// midRanks ranks values from 1, averaging the ranks within a tied group, and
|
||||||
|
// also returns the size of every group of size two or more.
|
||||||
|
func midRanks(values []float64) ([]float64, []int) {
|
||||||
|
order := make([]int, len(values))
|
||||||
|
for index := range order {
|
||||||
|
order[index] = index
|
||||||
|
}
|
||||||
|
slices.SortStableFunc(order, func(left, right int) int {
|
||||||
|
switch {
|
||||||
|
case values[left] < values[right]:
|
||||||
|
return -1
|
||||||
|
case values[left] > values[right]:
|
||||||
|
return 1
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
})
|
||||||
|
ranks := make([]float64, len(values))
|
||||||
|
var tieGroups []int
|
||||||
|
for start := 0; start < len(order); {
|
||||||
|
end := start + 1
|
||||||
|
for end < len(order) && values[order[end]] == values[order[start]] {
|
||||||
|
end++
|
||||||
|
}
|
||||||
|
shared := float64(start+1+end) / 2
|
||||||
|
for index := start; index < end; index++ {
|
||||||
|
ranks[order[index]] = shared
|
||||||
|
}
|
||||||
|
if end-start > 1 {
|
||||||
|
tieGroups = append(tieGroups, end-start)
|
||||||
|
}
|
||||||
|
start = end
|
||||||
|
}
|
||||||
|
return ranks, tieGroups
|
||||||
|
}
|
||||||
@@ -0,0 +1,211 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// R's wilcox.test on the Hollander and Wolfe (1973), 69f chorioamnion data
|
||||||
|
// reports W = 35 with an exact two-sided p-value of 0.2544; the one-sided
|
||||||
|
// greater alternative that the help page uses reports the same W with
|
||||||
|
// p-value = 0.1272.
|
||||||
|
func TestRankSum_MatchesPublishedChorioamnionResult(t *testing.T) {
|
||||||
|
result := rankSum(chorioamnionTerm, chorioamnionEarly)
|
||||||
|
|
||||||
|
if result.Statistic != 35 {
|
||||||
|
t.Errorf("statistic %v, want 35", result.Statistic)
|
||||||
|
}
|
||||||
|
if !result.Exact {
|
||||||
|
t.Error("expected the exact null distribution for untied samples this small")
|
||||||
|
}
|
||||||
|
if math.Abs(result.PValue-0.2544) > 5e-5 {
|
||||||
|
t.Errorf("two-sided p-value %.6f, want 0.2544", result.PValue)
|
||||||
|
}
|
||||||
|
upper := exactRankSumUpperTail(35, len(chorioamnionTerm), len(chorioamnionEarly))
|
||||||
|
if math.Abs(upper-0.1272) > 5e-5 {
|
||||||
|
t.Errorf("one-sided p-value %.6f, want 0.1272", upper)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A12 is P(X > Y) + 0.5 P(X = Y), which is U/(mn). With the published W = 35
|
||||||
|
// and sample sizes 10 and 5 the effect size is 35/50 = 0.70. The expected value
|
||||||
|
// is therefore the published Mann-Whitney statistic combined with the published
|
||||||
|
// definition in Vargha and Delaney (2000), not a number this tool produced.
|
||||||
|
func TestVargaDelaneyA12_MatchesPublishedChorioamnionStatistic(t *testing.T) {
|
||||||
|
got := vargaDelaneyA12(chorioamnionTerm, chorioamnionEarly)
|
||||||
|
if math.Abs(got-0.70) > 1e-12 {
|
||||||
|
t.Errorf("A12 = %v, want 0.70", got)
|
||||||
|
}
|
||||||
|
if reversed := vargaDelaneyA12(chorioamnionEarly, chorioamnionTerm); math.Abs(reversed-0.30) > 1e-12 {
|
||||||
|
t.Errorf("reversed A12 = %v, want 0.30", reversed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Identical samples are stochastically equal, which Vargha and Delaney define
|
||||||
|
// as A = 0.5, and complete separation gives 1 and 0.
|
||||||
|
func TestVargaDelaneyA12_BoundaryCases(t *testing.T) {
|
||||||
|
same := []float64{1, 2, 3, 4}
|
||||||
|
if got := vargaDelaneyA12(same, same); got != 0.5 {
|
||||||
|
t.Errorf("A12 of a sample against itself = %v, want 0.5", got)
|
||||||
|
}
|
||||||
|
if got := vargaDelaneyA12([]float64{5, 6, 7}, []float64{1, 2}); got != 1 {
|
||||||
|
t.Errorf("A12 with complete dominance = %v, want 1", got)
|
||||||
|
}
|
||||||
|
if got := vargaDelaneyA12([]float64{1, 2}, []float64{5, 6, 7}); got != 0 {
|
||||||
|
t.Errorf("A12 with complete subordination = %v, want 0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The counting definition and the rank-sum route must agree, including when the
|
||||||
|
// samples are tied against each other, which is the case the evaluation data is
|
||||||
|
// always in because censored runs are all held at the budget.
|
||||||
|
func TestVargaDelaneyA12_AgreesWithRankSumStatistic(t *testing.T) {
|
||||||
|
cases := [][2][]float64{
|
||||||
|
{{1, 2, 3}, {2, 3, 4}},
|
||||||
|
{{40, 40, 40, 12}, {40, 7, 3}},
|
||||||
|
{{5}, {5, 5, 5}},
|
||||||
|
{{9, 9, 9}, {9, 9, 9}},
|
||||||
|
}
|
||||||
|
for _, test := range cases {
|
||||||
|
result := rankSum(test[0], test[1])
|
||||||
|
expected := result.Statistic / float64(len(test[0])*len(test[1]))
|
||||||
|
if math.Abs(result.A12-expected) > 1e-12 {
|
||||||
|
t.Errorf("A12 %v for %v vs %v, want U/(mn) = %v", result.A12, test[0], test[1], expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The tie-corrected variance is checked against the exact permutation variance
|
||||||
|
// of the statistic, computed here by enumerating every way to split the pooled
|
||||||
|
// midranks. That is an independent calculation, not a second call into the
|
||||||
|
// implementation under test.
|
||||||
|
func TestRankSumVariance_MatchesExactPermutationVariance(t *testing.T) {
|
||||||
|
cases := [][]float64{
|
||||||
|
{1, 2, 3, 4, 5, 6, 7, 8},
|
||||||
|
{40, 40, 40, 40, 12, 7, 3, 3},
|
||||||
|
{5, 5, 5, 5, 5, 5, 5, 9},
|
||||||
|
{2, 2, 3, 3, 3, 4, 9, 9, 9},
|
||||||
|
}
|
||||||
|
for _, pooled := range cases {
|
||||||
|
firstSize := len(pooled) / 2
|
||||||
|
ranks, tieGroups := midRanks(pooled)
|
||||||
|
mean, variance := permutationMomentsOfRankSum(ranks, firstSize)
|
||||||
|
expectedMean := float64(firstSize*(len(pooled)-firstSize)) / 2
|
||||||
|
if math.Abs(mean-expectedMean) > 1e-9 {
|
||||||
|
t.Errorf("%v: permutation mean %v, want %v", pooled, mean, expectedMean)
|
||||||
|
}
|
||||||
|
got := rankSumVariance(firstSize, len(pooled)-firstSize, tieGroups)
|
||||||
|
if math.Abs(got-variance) > 1e-9 {
|
||||||
|
t.Errorf("%v: tie-corrected variance %v, want the permutation variance %v", pooled, got, variance)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// permutationMomentsOfRankSum enumerates every subset of the given size and
|
||||||
|
// returns the mean and variance of the Mann-Whitney statistic over them.
|
||||||
|
func permutationMomentsOfRankSum(ranks []float64, firstSize int) (float64, float64) {
|
||||||
|
offset := float64(firstSize) * float64(firstSize+1) / 2
|
||||||
|
var values []float64
|
||||||
|
chosen := make([]int, 0, firstSize)
|
||||||
|
var walk func(start int)
|
||||||
|
walk = func(start int) {
|
||||||
|
if len(chosen) == firstSize {
|
||||||
|
total := 0.0
|
||||||
|
for _, index := range chosen {
|
||||||
|
total += ranks[index]
|
||||||
|
}
|
||||||
|
values = append(values, total-offset)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for index := start; index < len(ranks); index++ {
|
||||||
|
chosen = append(chosen, index)
|
||||||
|
walk(index + 1)
|
||||||
|
chosen = chosen[:len(chosen)-1]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
walk(0)
|
||||||
|
mean := 0.0
|
||||||
|
for _, value := range values {
|
||||||
|
mean += value
|
||||||
|
}
|
||||||
|
mean /= float64(len(values))
|
||||||
|
variance := 0.0
|
||||||
|
for _, value := range values {
|
||||||
|
variance += (value - mean) * (value - mean)
|
||||||
|
}
|
||||||
|
return mean, variance / float64(len(values))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMidRanks_AveragesTiedGroups(t *testing.T) {
|
||||||
|
ranks, tieGroups := midRanks([]float64{3, 1, 3, 2, 3})
|
||||||
|
expected := []float64{4, 1, 4, 2, 4}
|
||||||
|
for index, want := range expected {
|
||||||
|
if ranks[index] != want {
|
||||||
|
t.Errorf("rank %d = %v, want %v", index, ranks[index], want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(tieGroups) != 1 || tieGroups[0] != 3 {
|
||||||
|
t.Errorf("tie groups %v, want [3]", tieGroups)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRankSum_TiedSamplesUseTheNormalApproximation(t *testing.T) {
|
||||||
|
result := rankSum([]float64{1, 2, 3, 4}, []float64{3, 4, 5, 6})
|
||||||
|
if result.Exact {
|
||||||
|
t.Error("used the exact null distribution despite ties")
|
||||||
|
}
|
||||||
|
if math.IsNaN(result.PValue) || result.PValue < 0 || result.PValue > 1 {
|
||||||
|
t.Errorf("p-value %v", result.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Every observation identical carries no information, and the test must say so
|
||||||
|
// rather than dividing by a zero variance.
|
||||||
|
func TestRankSum_AllValuesIdentical(t *testing.T) {
|
||||||
|
result := rankSum([]float64{40, 40, 40}, []float64{40, 40, 40, 40})
|
||||||
|
if result.PValue != 1 {
|
||||||
|
t.Errorf("p-value %v, want 1", result.PValue)
|
||||||
|
}
|
||||||
|
if result.A12 != 0.5 {
|
||||||
|
t.Errorf("A12 %v, want 0.5", result.A12)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRankSum_SingleObservationPerSample(t *testing.T) {
|
||||||
|
result := rankSum([]float64{3}, []float64{9})
|
||||||
|
if result.Statistic != 0 {
|
||||||
|
t.Errorf("statistic %v, want 0", result.Statistic)
|
||||||
|
}
|
||||||
|
if result.A12 != 0 {
|
||||||
|
t.Errorf("A12 %v, want 0", result.A12)
|
||||||
|
}
|
||||||
|
if math.IsNaN(result.PValue) || result.PValue > 1 {
|
||||||
|
t.Errorf("p-value %v", result.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRankSum_EmptySampleHasNoStatistic(t *testing.T) {
|
||||||
|
result := rankSum(nil, []float64{1, 2, 3})
|
||||||
|
if !math.IsNaN(result.PValue) || !math.IsNaN(result.A12) {
|
||||||
|
t.Errorf("result %+v, want everything undefined", result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The exact null distribution must be a proper distribution: the counts sum to
|
||||||
|
// the binomial coefficient and the distribution is symmetric about mn/2.
|
||||||
|
func TestExactRankSumCounts_FormAProperSymmetricDistribution(t *testing.T) {
|
||||||
|
counts := exactRankSumCounts(4, 6)
|
||||||
|
total := 0.0
|
||||||
|
for _, count := range counts {
|
||||||
|
total += count
|
||||||
|
}
|
||||||
|
if total != 210 {
|
||||||
|
t.Errorf("counts sum to %v, want C(10,4) = 210", total)
|
||||||
|
}
|
||||||
|
for index := range counts {
|
||||||
|
mirrored := counts[len(counts)-1-index]
|
||||||
|
if counts[index] != mirrored {
|
||||||
|
t.Errorf("count at %d is %v but %v at the mirrored point", index, counts[index], mirrored)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,154 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"maps"
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"text/tabwriter"
|
||||||
|
)
|
||||||
|
|
||||||
|
func writeReport(result analysis, out io.Writer) {
|
||||||
|
fmt.Fprintf(out, "primary outcome: %s\n\n", result.Outcome)
|
||||||
|
|
||||||
|
writeTable(out, []string{"arm", "runs", "violated", "censored", "excluded", "missing", "median steps", "violation rate"},
|
||||||
|
func(add func(...string)) {
|
||||||
|
for _, summary := range result.Arms {
|
||||||
|
add(
|
||||||
|
summary.Arm,
|
||||||
|
strconv.Itoa(summary.Usable),
|
||||||
|
strconv.Itoa(summary.Violated),
|
||||||
|
strconv.Itoa(summary.Censored),
|
||||||
|
strconv.Itoa(summary.Excluded),
|
||||||
|
strconv.Itoa(len(summary.MissingSeeds)),
|
||||||
|
formatMedian(summary.MedianStepsToFirstViolation),
|
||||||
|
formatRatio(summary.ViolationRate, 3),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
fmt.Fprintln(out)
|
||||||
|
fmt.Fprintln(out, "a detection is one distinct property violated in one run; run hours sum the per-run wall clock")
|
||||||
|
writeTable(out, []string{"arm", "actions", "run hours", "detections", "defects/1k actions", "defects/hour", "distinct defects", "found in one run"},
|
||||||
|
func(add func(...string)) {
|
||||||
|
for _, summary := range result.Arms {
|
||||||
|
add(
|
||||||
|
summary.Arm,
|
||||||
|
strconv.Itoa(summary.TotalActions),
|
||||||
|
fmt.Sprintf("%.2f", summary.TotalRunHours),
|
||||||
|
strconv.Itoa(summary.Detections),
|
||||||
|
formatRatio(summary.DefectsPerThousandActions, 2),
|
||||||
|
formatRatio(summary.DefectsPerHour, 2),
|
||||||
|
strconv.Itoa(summary.DistinctDefects),
|
||||||
|
formatSingletons(summary),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, summary := range result.Arms {
|
||||||
|
if len(summary.ExcludedByReason) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var parts []string
|
||||||
|
for _, reason := range sortedKeys(summary.ExcludedByReason) {
|
||||||
|
parts = append(parts, fmt.Sprintf("%s=%d", reason, summary.ExcludedByReason[reason]))
|
||||||
|
}
|
||||||
|
fmt.Fprintf(out, "\n%s excluded %d run(s) as missing data, not as censored observations: %s",
|
||||||
|
summary.Arm, summary.Excluded, strings.Join(parts, ", "))
|
||||||
|
}
|
||||||
|
for _, summary := range result.Arms {
|
||||||
|
if summary.EventsHeldAtBudget > 0 {
|
||||||
|
fmt.Fprintf(out, "\n%s held %d violation(s) reported past the budget at %d steps",
|
||||||
|
summary.Arm, summary.EventsHeldAtBudget, summary.StepBudget)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(result.Arms) > 0 {
|
||||||
|
fmt.Fprintln(out)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.LogRank != nil {
|
||||||
|
fmt.Fprintf(out, "\nlog-rank across %d arms: chi-square %.4f on %d df, p %s\n",
|
||||||
|
len(result.LogRank.Groups), result.LogRank.ChiSquare, result.LogRank.DegreesOfFreedom,
|
||||||
|
formatPValue(result.LogRank.PValue))
|
||||||
|
writeTable(out, []string{"arm", "n", "observed", "expected"}, func(add func(...string)) {
|
||||||
|
for index, name := range result.LogRank.Groups {
|
||||||
|
add(name,
|
||||||
|
strconv.Itoa(result.LogRank.Sizes[index]),
|
||||||
|
fmt.Sprintf("%.0f", result.LogRank.Observed[index]),
|
||||||
|
fmt.Sprintf("%.2f", result.LogRank.Expected[index]))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(result.Pairwise) > 0 {
|
||||||
|
fmt.Fprintln(out, "\npairwise wilcoxon rank-sum, censored runs held at the budget")
|
||||||
|
fmt.Fprintln(out, "a12 above 0.5 means the first arm takes more steps to its first violation")
|
||||||
|
writeTable(out, []string{"comparison", "n1", "n2", "u", "a12", "p", "holm p"}, func(add func(...string)) {
|
||||||
|
for _, pair := range result.Pairwise {
|
||||||
|
add(
|
||||||
|
pair.First+" vs "+pair.Second,
|
||||||
|
strconv.Itoa(pair.FirstSize),
|
||||||
|
strconv.Itoa(pair.SecondSize),
|
||||||
|
fmt.Sprintf("%.1f", pair.Statistic),
|
||||||
|
fmt.Sprintf("%.3f", pair.A12),
|
||||||
|
formatPValue(pair.PValue),
|
||||||
|
formatPValue(pair.HolmPValue),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, note := range result.Notes {
|
||||||
|
fmt.Fprintf(out, "\nnote: %s\n", note)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedKeys(counts map[string]int) []string {
|
||||||
|
return slices.Sorted(maps.Keys(counts))
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeTable(out io.Writer, header []string, rows func(add func(...string))) {
|
||||||
|
writer := tabwriter.NewWriter(out, 0, 0, 2, ' ', 0)
|
||||||
|
fmt.Fprintln(writer, strings.Join(header, "\t"))
|
||||||
|
rows(func(cells ...string) {
|
||||||
|
fmt.Fprintln(writer, strings.Join(cells, "\t"))
|
||||||
|
})
|
||||||
|
writer.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
// formatMedian says undefined rather than substituting a mean, because a curve
|
||||||
|
// that never reaches one half has no median to report.
|
||||||
|
func formatMedian(value *float64) string {
|
||||||
|
if value == nil {
|
||||||
|
return "undefined"
|
||||||
|
}
|
||||||
|
return strconv.FormatFloat(*value, 'f', -1, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatRatio(value *float64, digits int) string {
|
||||||
|
if value == nil {
|
||||||
|
return "n/a"
|
||||||
|
}
|
||||||
|
return strconv.FormatFloat(*value, 'f', digits, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatSingletons(summary armSummary) string {
|
||||||
|
if summary.SingletonFraction == nil {
|
||||||
|
return "n/a"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%d/%d (%.3f)", summary.SingletonDefects, summary.DistinctDefects, *summary.SingletonFraction)
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatPValue(value float64) string {
|
||||||
|
switch {
|
||||||
|
case math.IsNaN(value):
|
||||||
|
return "n/a"
|
||||||
|
case value < 1e-4:
|
||||||
|
return fmt.Sprintf("%.3e", value)
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("%.4f", value)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
|
)
|
||||||
|
|
||||||
|
// observation is one run reduced to what the survival analysis needs: the step
|
||||||
|
// count at which it left the risk set, and whether it left because a violation
|
||||||
|
// was found (an event) or because it exhausted the step budget without one
|
||||||
|
// (right-censored). A run that failed or timed out is neither and never
|
||||||
|
// reaches this type.
|
||||||
|
type observation struct {
|
||||||
|
Steps float64
|
||||||
|
Event bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type survivalPoint struct {
|
||||||
|
Steps float64 `json:"steps"`
|
||||||
|
AtRisk int `json:"at_risk"`
|
||||||
|
Events int `json:"events"`
|
||||||
|
Censored int `json:"censored"`
|
||||||
|
Survival float64 `json:"survival"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// kaplanMeier is the product-limit estimate of Kaplan and Meier (1958), one row
|
||||||
|
// per distinct observed step count. Runs censored at a step count tied with an
|
||||||
|
// event are counted in the risk set for that event, which is the standard
|
||||||
|
// convention.
|
||||||
|
func kaplanMeier(observations []observation) []survivalPoint {
|
||||||
|
if len(observations) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
remaining := len(observations)
|
||||||
|
survival := 1.0
|
||||||
|
var curve []survivalPoint
|
||||||
|
for _, steps := range distinctSteps(observations) {
|
||||||
|
events, censored := 0, 0
|
||||||
|
for _, item := range observations {
|
||||||
|
if item.Steps != steps {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if item.Event {
|
||||||
|
events++
|
||||||
|
} else {
|
||||||
|
censored++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
atRisk := remaining
|
||||||
|
if events > 0 {
|
||||||
|
survival *= 1 - float64(events)/float64(atRisk)
|
||||||
|
}
|
||||||
|
curve = append(curve, survivalPoint{
|
||||||
|
Steps: steps,
|
||||||
|
AtRisk: atRisk,
|
||||||
|
Events: events,
|
||||||
|
Censored: censored,
|
||||||
|
Survival: survival,
|
||||||
|
})
|
||||||
|
remaining -= events + censored
|
||||||
|
}
|
||||||
|
return curve
|
||||||
|
}
|
||||||
|
|
||||||
|
// medianSurvival is the smallest step count at which the estimate falls to or
|
||||||
|
// below one half. It is undefined whenever fewer than half the runs violate,
|
||||||
|
// and the second return value says so: substituting a mean there would report a
|
||||||
|
// number the data does not contain.
|
||||||
|
func medianSurvival(curve []survivalPoint) (float64, bool) {
|
||||||
|
for _, point := range curve {
|
||||||
|
if point.Survival <= 0.5 {
|
||||||
|
return point.Steps, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func distinctSteps(observations []observation) []float64 {
|
||||||
|
steps := make([]float64, 0, len(observations))
|
||||||
|
for _, item := range observations {
|
||||||
|
steps = append(steps, item.Steps)
|
||||||
|
}
|
||||||
|
slices.Sort(steps)
|
||||||
|
return slices.Compact(steps)
|
||||||
|
}
|
||||||
|
|
||||||
|
type logRankResult struct {
|
||||||
|
Groups []string `json:"groups"`
|
||||||
|
Sizes []int `json:"sizes"`
|
||||||
|
Observed []float64 `json:"observed"`
|
||||||
|
Expected []float64 `json:"expected"`
|
||||||
|
ChiSquare float64 `json:"chi_square"`
|
||||||
|
DegreesOfFreedom int `json:"degrees_of_freedom"`
|
||||||
|
PValue float64 `json:"p_value"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// logRank is the k-sample Mantel-Haenszel log-rank test. At every distinct
|
||||||
|
// event time it contrasts observed with expected events under the null of equal
|
||||||
|
// hazards, then combines the k-1 independent differences through their
|
||||||
|
// covariance matrix: chi-square = U' V^-1 U on k-1 degrees of freedom.
|
||||||
|
// Mantel (1966); Peto and Peto (1972); Klein and Moeschberger, Survival
|
||||||
|
// Analysis, 2nd ed., section 7.3.
|
||||||
|
func logRank(names []string, groups [][]observation) logRankResult {
|
||||||
|
var keptNames []string
|
||||||
|
var kept [][]observation
|
||||||
|
for index, group := range groups {
|
||||||
|
if len(group) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
keptNames = append(keptNames, names[index])
|
||||||
|
kept = append(kept, group)
|
||||||
|
}
|
||||||
|
names, groups = keptNames, kept
|
||||||
|
|
||||||
|
count := len(groups)
|
||||||
|
result := logRankResult{
|
||||||
|
Groups: names,
|
||||||
|
Sizes: make([]int, count),
|
||||||
|
Observed: make([]float64, count),
|
||||||
|
Expected: make([]float64, count),
|
||||||
|
DegreesOfFreedom: count - 1,
|
||||||
|
PValue: math.NaN(),
|
||||||
|
}
|
||||||
|
if count < 2 {
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
var pooled []observation
|
||||||
|
for index, group := range groups {
|
||||||
|
result.Sizes[index] = len(group)
|
||||||
|
pooled = append(pooled, group...)
|
||||||
|
}
|
||||||
|
|
||||||
|
covariance := make([][]float64, count)
|
||||||
|
for index := range covariance {
|
||||||
|
covariance[index] = make([]float64, count)
|
||||||
|
}
|
||||||
|
atRisk := make([]float64, count)
|
||||||
|
deaths := make([]float64, count)
|
||||||
|
|
||||||
|
for _, steps := range distinctSteps(pooled) {
|
||||||
|
totalAtRisk, totalDeaths := 0.0, 0.0
|
||||||
|
for index, group := range groups {
|
||||||
|
atRisk[index], deaths[index] = 0, 0
|
||||||
|
for _, item := range group {
|
||||||
|
if item.Steps >= steps {
|
||||||
|
atRisk[index]++
|
||||||
|
}
|
||||||
|
if item.Steps == steps && item.Event {
|
||||||
|
deaths[index]++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
totalAtRisk += atRisk[index]
|
||||||
|
totalDeaths += deaths[index]
|
||||||
|
}
|
||||||
|
if totalDeaths == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for index := range groups {
|
||||||
|
result.Observed[index] += deaths[index]
|
||||||
|
result.Expected[index] += totalDeaths * atRisk[index] / totalAtRisk
|
||||||
|
}
|
||||||
|
if totalAtRisk <= 1 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
scale := totalDeaths * (totalAtRisk - totalDeaths) / (totalAtRisk - 1)
|
||||||
|
for row := range groups {
|
||||||
|
share := atRisk[row] / totalAtRisk
|
||||||
|
covariance[row][row] += scale * share * (1 - share)
|
||||||
|
for column := range groups {
|
||||||
|
if column == row {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
covariance[row][column] -= scale * share * atRisk[column] / totalAtRisk
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
reduced := make([][]float64, count-1)
|
||||||
|
difference := make([]float64, count-1)
|
||||||
|
for row := 0; row < count-1; row++ {
|
||||||
|
reduced[row] = make([]float64, count-1)
|
||||||
|
copy(reduced[row], covariance[row][:count-1])
|
||||||
|
difference[row] = result.Observed[row] - result.Expected[row]
|
||||||
|
}
|
||||||
|
solution, ok := solveLinearSystem(reduced, difference)
|
||||||
|
if !ok {
|
||||||
|
result.ChiSquare = 0
|
||||||
|
result.PValue = 1
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
statistic := 0.0
|
||||||
|
for index := range difference {
|
||||||
|
statistic += difference[index] * solution[index]
|
||||||
|
}
|
||||||
|
if statistic < 0 || math.IsNaN(statistic) {
|
||||||
|
statistic = 0
|
||||||
|
}
|
||||||
|
result.ChiSquare = statistic
|
||||||
|
result.PValue = chiSquareUpperTail(statistic, result.DegreesOfFreedom)
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// solveLinearSystem solves matrix*x = vector by Gaussian elimination with
|
||||||
|
// partial pivoting, reporting failure rather than a value when the matrix is
|
||||||
|
// singular, which is what a group with no events at all produces.
|
||||||
|
func solveLinearSystem(matrix [][]float64, vector []float64) ([]float64, bool) {
|
||||||
|
size := len(vector)
|
||||||
|
work := make([][]float64, size)
|
||||||
|
for row := range work {
|
||||||
|
work[row] = make([]float64, size+1)
|
||||||
|
copy(work[row], matrix[row])
|
||||||
|
work[row][size] = vector[row]
|
||||||
|
}
|
||||||
|
for column := 0; column < size; column++ {
|
||||||
|
pivot := column
|
||||||
|
for row := column + 1; row < size; row++ {
|
||||||
|
if math.Abs(work[row][column]) > math.Abs(work[pivot][column]) {
|
||||||
|
pivot = row
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if math.Abs(work[pivot][column]) < 1e-12 {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
work[column], work[pivot] = work[pivot], work[column]
|
||||||
|
for row := column + 1; row < size; row++ {
|
||||||
|
factor := work[row][column] / work[column][column]
|
||||||
|
for next := column; next <= size; next++ {
|
||||||
|
work[row][next] -= factor * work[column][next]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
solution := make([]float64, size)
|
||||||
|
for row := size - 1; row >= 0; row-- {
|
||||||
|
total := work[row][size]
|
||||||
|
for column := row + 1; column < size; column++ {
|
||||||
|
total -= work[row][column] * solution[column]
|
||||||
|
}
|
||||||
|
solution[row] = total / work[row][row]
|
||||||
|
}
|
||||||
|
return solution, true
|
||||||
|
}
|
||||||
@@ -0,0 +1,304 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"math/rand"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The product-limit estimates for the 6-MP arm of Freireich et al. (1963) are
|
||||||
|
// the worked example reproduced in Collett, Modelling Survival Data in Medical
|
||||||
|
// Research, and in the standard course treatments of the gehan data:
|
||||||
|
//
|
||||||
|
// t: 6 7 10 13 16 22 23
|
||||||
|
// S(t): 0.857 0.807 0.753 0.690 0.627 0.538 0.448
|
||||||
|
func TestKaplanMeier_MatchesPublishedGehanEstimates(t *testing.T) {
|
||||||
|
curve := kaplanMeier(gehanSixMercaptopurine)
|
||||||
|
expected := map[float64]float64{
|
||||||
|
6: 0.857, 7: 0.807, 10: 0.753, 13: 0.690, 16: 0.627, 22: 0.538, 23: 0.448,
|
||||||
|
}
|
||||||
|
seen := 0
|
||||||
|
for _, point := range curve {
|
||||||
|
want, ok := expected[point.Steps]
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen++
|
||||||
|
if math.Abs(point.Survival-want) > 5e-4 {
|
||||||
|
t.Errorf("S(%v) = %.4f, want %v", point.Steps, point.Survival, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if seen != len(expected) {
|
||||||
|
t.Fatalf("matched %d of %d published times", seen, len(expected))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The risk set at each time is the count of runs still under observation, with
|
||||||
|
// runs censored at a tied time counted as at risk for that event.
|
||||||
|
func TestKaplanMeier_RiskSetHandlesTiesAndCensoring(t *testing.T) {
|
||||||
|
curve := kaplanMeier(gehanSixMercaptopurine)
|
||||||
|
expected := map[float64]struct {
|
||||||
|
atRisk int
|
||||||
|
events int
|
||||||
|
censored int
|
||||||
|
}{
|
||||||
|
6: {21, 3, 1},
|
||||||
|
7: {17, 1, 0},
|
||||||
|
9: {16, 0, 1},
|
||||||
|
10: {15, 1, 1},
|
||||||
|
13: {12, 1, 0},
|
||||||
|
23: {6, 1, 0},
|
||||||
|
}
|
||||||
|
for _, point := range curve {
|
||||||
|
want, ok := expected[point.Steps]
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if point.AtRisk != want.atRisk || point.Events != want.events || point.Censored != want.censored {
|
||||||
|
t.Errorf("at %v: risk=%d events=%d censored=%d, want risk=%d events=%d censored=%d",
|
||||||
|
point.Steps, point.AtRisk, point.Events, point.Censored, want.atRisk, want.events, want.censored)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Published medians: 23 weeks for 6-MP against 8 weeks for placebo (Gehan and
|
||||||
|
// Freireich, "The 6-MP versus placebo clinical trial in acute leukemia",
|
||||||
|
// Clinical Trials 8(3), 2011), and 31 against 23 weeks for the aml arms as
|
||||||
|
// reported by survfit in R's survival package.
|
||||||
|
func TestMedianSurvival_MatchesPublishedMedians(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
observations []observation
|
||||||
|
expected float64
|
||||||
|
}{
|
||||||
|
{"gehan 6-MP", gehanSixMercaptopurine, 23},
|
||||||
|
{"gehan placebo", gehanPlacebo, 8},
|
||||||
|
{"aml maintained", amlMaintained, 31},
|
||||||
|
{"aml nonmaintained", amlNonmaintained, 23},
|
||||||
|
}
|
||||||
|
for _, test := range cases {
|
||||||
|
median, ok := medianSurvival(kaplanMeier(test.observations))
|
||||||
|
if !ok {
|
||||||
|
t.Errorf("%s: median undefined, want %v", test.name, test.expected)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if median != test.expected {
|
||||||
|
t.Errorf("%s: median %v, want %v", test.name, median, test.expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// R's survival package documents this log-rank on the aml data:
|
||||||
|
//
|
||||||
|
// N Observed Expected (O-E)^2/E (O-E)^2/V
|
||||||
|
// x=Maintained 11 7 10.69 1.27 3.4
|
||||||
|
// x=Nonmaintained 12 11 7.31 1.86 3.4
|
||||||
|
// Chisq= 3.4 on 1 degrees of freedom, p= 0.0653
|
||||||
|
func TestLogRank_MatchesPublishedAmlResult(t *testing.T) {
|
||||||
|
result := logRank([]string{"maintained", "nonmaintained"}, [][]observation{amlMaintained, amlNonmaintained})
|
||||||
|
|
||||||
|
if result.Observed[0] != 7 || result.Observed[1] != 11 {
|
||||||
|
t.Errorf("observed %v, want [7 11]", result.Observed)
|
||||||
|
}
|
||||||
|
if math.Abs(result.Expected[0]-10.69) > 5e-3 || math.Abs(result.Expected[1]-7.31) > 5e-3 {
|
||||||
|
t.Errorf("expected %v, want [10.69 7.31]", result.Expected)
|
||||||
|
}
|
||||||
|
for index, want := range []float64{1.27, 1.86} {
|
||||||
|
difference := result.Observed[index] - result.Expected[index]
|
||||||
|
got := difference * difference / result.Expected[index]
|
||||||
|
if math.Abs(got-want) > 5e-3 {
|
||||||
|
t.Errorf("(O-E)^2/E for group %d = %.4f, want %v", index, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if math.Abs(result.ChiSquare-3.4) > 5e-2 {
|
||||||
|
t.Errorf("chi-square %.4f, want 3.4", result.ChiSquare)
|
||||||
|
}
|
||||||
|
if math.Abs(result.PValue-0.0653) > 5e-4 {
|
||||||
|
t.Errorf("p-value %.6f, want 0.0653", result.PValue)
|
||||||
|
}
|
||||||
|
if result.DegreesOfFreedom != 1 {
|
||||||
|
t.Errorf("degrees of freedom %d, want 1", result.DegreesOfFreedom)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The log-rank on the Freireich 6-MP trial is the textbook worked example:
|
||||||
|
// observed 9 against 19.25 expected in the treated arm and 21 against 10.75 in
|
||||||
|
// the control arm, Mantel-Haenszel chi-square 16.79 on 1 degree of freedom,
|
||||||
|
// p = 4.17e-05. Reported for instance in Rodriguez, Kaplan-Meier and
|
||||||
|
// Mantel-Haenszel, https://grodri.github.io/survival/gehan, and in Collett.
|
||||||
|
func TestLogRank_MatchesPublishedGehanResult(t *testing.T) {
|
||||||
|
result := logRank([]string{"6-MP", "placebo"}, [][]observation{gehanSixMercaptopurine, gehanPlacebo})
|
||||||
|
|
||||||
|
if result.Observed[0] != 9 || result.Observed[1] != 21 {
|
||||||
|
t.Errorf("observed %v, want [9 21]", result.Observed)
|
||||||
|
}
|
||||||
|
if math.Abs(result.Expected[0]-19.25) > 5e-3 || math.Abs(result.Expected[1]-10.75) > 5e-3 {
|
||||||
|
t.Errorf("expected %v, want [19.25 10.75]", result.Expected)
|
||||||
|
}
|
||||||
|
if math.Abs(result.ChiSquare-16.79) > 5e-3 {
|
||||||
|
t.Errorf("chi-square %.4f, want 16.79", result.ChiSquare)
|
||||||
|
}
|
||||||
|
if math.Abs(result.PValue-4.17e-5) > 5e-8 {
|
||||||
|
t.Errorf("p-value %.3e, want 4.17e-05", result.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// An arm with no usable runs contributes nothing and must not consume a degree
|
||||||
|
// of freedom or make the covariance matrix singular.
|
||||||
|
func TestLogRank_EmptyGroupIsDropped(t *testing.T) {
|
||||||
|
two := logRank([]string{"a", "b"}, [][]observation{amlMaintained, amlNonmaintained})
|
||||||
|
three := logRank([]string{"a", "b", "c"}, [][]observation{amlMaintained, amlNonmaintained, nil})
|
||||||
|
if math.Abs(two.ChiSquare-three.ChiSquare) > 1e-12 {
|
||||||
|
t.Errorf("chi-square %.10f with an empty third group, want %.10f", three.ChiSquare, two.ChiSquare)
|
||||||
|
}
|
||||||
|
if three.DegreesOfFreedom != 1 {
|
||||||
|
t.Errorf("degrees of freedom %d, want 1", three.DegreesOfFreedom)
|
||||||
|
}
|
||||||
|
if len(three.Groups) != 2 {
|
||||||
|
t.Errorf("groups %v, want the empty arm dropped", three.Groups)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// No published multi-arm dataset with a printed log-rank chi-square was found
|
||||||
|
// small enough to embed, so the k-group covariance algebra is checked against
|
||||||
|
// its own null distribution instead: under the null of equal hazards the
|
||||||
|
// statistic is asymptotically chi-square on k-1 degrees of freedom, so its mean
|
||||||
|
// over random relabellings has to sit near k-1. A wrong variance term or a wrong
|
||||||
|
// degrees-of-freedom count moves this badly.
|
||||||
|
func TestLogRank_NullMeanTracksDegreesOfFreedom(t *testing.T) {
|
||||||
|
pooled := make([]observation, 0, 36)
|
||||||
|
for index := 0; index < 36; index++ {
|
||||||
|
pooled = append(pooled, observation{Steps: float64(index%17 + 1), Event: index%5 != 0})
|
||||||
|
}
|
||||||
|
for _, groupCount := range []int{2, 3, 4} {
|
||||||
|
generator := rand.New(rand.NewSource(20260812))
|
||||||
|
total := 0.0
|
||||||
|
const replicates = 4000
|
||||||
|
for replicate := 0; replicate < replicates; replicate++ {
|
||||||
|
shuffled := slices.Clone(pooled)
|
||||||
|
generator.Shuffle(len(shuffled), func(left, right int) {
|
||||||
|
shuffled[left], shuffled[right] = shuffled[right], shuffled[left]
|
||||||
|
})
|
||||||
|
names := make([]string, groupCount)
|
||||||
|
groups := make([][]observation, groupCount)
|
||||||
|
for index, item := range shuffled {
|
||||||
|
groups[index%groupCount] = append(groups[index%groupCount], item)
|
||||||
|
}
|
||||||
|
for index := range names {
|
||||||
|
names[index] = strconv.Itoa(index)
|
||||||
|
}
|
||||||
|
total += logRank(names, groups).ChiSquare
|
||||||
|
}
|
||||||
|
mean := total / replicates
|
||||||
|
expected := float64(groupCount - 1)
|
||||||
|
if math.Abs(mean-expected) > 0.15*expected {
|
||||||
|
t.Errorf("%d groups: null mean chi-square %.3f, want near %v", groupCount, mean, expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A three-group split of one homogeneous sample must not look significant, and
|
||||||
|
// the statistic must be finite on 2 degrees of freedom.
|
||||||
|
func TestLogRank_ThreeIdenticalGroupsAreNotSignificant(t *testing.T) {
|
||||||
|
group := []observation{{4, true}, {7, true}, {9, false}, {12, true}, {20, false}}
|
||||||
|
result := logRank([]string{"a", "b", "c"}, [][]observation{group, group, group})
|
||||||
|
if result.DegreesOfFreedom != 2 {
|
||||||
|
t.Fatalf("degrees of freedom %d, want 2", result.DegreesOfFreedom)
|
||||||
|
}
|
||||||
|
if result.ChiSquare > 1e-9 {
|
||||||
|
t.Errorf("chi-square %.10f for three identical groups, want 0", result.ChiSquare)
|
||||||
|
}
|
||||||
|
if math.Abs(result.PValue-1) > 1e-9 {
|
||||||
|
t.Errorf("p-value %v, want 1", result.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKaplanMeier_EveryObservationCensored(t *testing.T) {
|
||||||
|
observations := []observation{{40, false}, {40, false}, {40, false}}
|
||||||
|
curve := kaplanMeier(observations)
|
||||||
|
if len(curve) != 1 {
|
||||||
|
t.Fatalf("curve has %d points, want 1", len(curve))
|
||||||
|
}
|
||||||
|
if curve[0].Survival != 1 || curve[0].Events != 0 || curve[0].Censored != 3 {
|
||||||
|
t.Errorf("point %+v, want survival 1 with 3 censored", curve[0])
|
||||||
|
}
|
||||||
|
if _, ok := medianSurvival(curve); ok {
|
||||||
|
t.Error("median defined for an arm where nothing violated")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKaplanMeier_EveryObservationAnEvent(t *testing.T) {
|
||||||
|
curve := kaplanMeier([]observation{{2, true}, {4, true}, {6, true}, {8, true}})
|
||||||
|
last := curve[len(curve)-1]
|
||||||
|
if last.Survival != 0 {
|
||||||
|
t.Errorf("final survival %v, want 0", last.Survival)
|
||||||
|
}
|
||||||
|
median, ok := medianSurvival(curve)
|
||||||
|
if !ok || median != 4 {
|
||||||
|
t.Errorf("median %v ok=%v, want 4", median, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKaplanMeier_TiedEventTimesDropOnce(t *testing.T) {
|
||||||
|
curve := kaplanMeier([]observation{{5, true}, {5, true}, {5, true}, {9, true}})
|
||||||
|
if len(curve) != 2 {
|
||||||
|
t.Fatalf("curve has %d points, want 2", len(curve))
|
||||||
|
}
|
||||||
|
if curve[0].Events != 3 || math.Abs(curve[0].Survival-0.25) > 1e-12 {
|
||||||
|
t.Errorf("first point %+v, want 3 events and survival 0.25", curve[0])
|
||||||
|
}
|
||||||
|
if curve[1].Survival != 0 {
|
||||||
|
t.Errorf("second point %+v, want survival 0", curve[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKaplanMeier_SingleObservation(t *testing.T) {
|
||||||
|
event := kaplanMeier([]observation{{11, true}})
|
||||||
|
if len(event) != 1 || event[0].Survival != 0 {
|
||||||
|
t.Fatalf("single event curve %+v", event)
|
||||||
|
}
|
||||||
|
median, ok := medianSurvival(event)
|
||||||
|
if !ok || median != 11 {
|
||||||
|
t.Errorf("median %v ok=%v, want 11", median, ok)
|
||||||
|
}
|
||||||
|
censored := kaplanMeier([]observation{{11, false}})
|
||||||
|
if len(censored) != 1 || censored[0].Survival != 1 {
|
||||||
|
t.Fatalf("single censored curve %+v", censored)
|
||||||
|
}
|
||||||
|
if _, ok := medianSurvival(censored); ok {
|
||||||
|
t.Error("median defined for a single censored observation")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKaplanMeier_NoObservations(t *testing.T) {
|
||||||
|
if curve := kaplanMeier(nil); curve != nil {
|
||||||
|
t.Errorf("curve %v for no observations, want nil", curve)
|
||||||
|
}
|
||||||
|
if _, ok := medianSurvival(nil); ok {
|
||||||
|
t.Error("median defined for an empty curve")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLogRank_SingleObservationPerGroup(t *testing.T) {
|
||||||
|
result := logRank([]string{"a", "b"}, [][]observation{{{3, true}}, {{9, true}}})
|
||||||
|
if math.IsNaN(result.ChiSquare) || result.ChiSquare < 0 {
|
||||||
|
t.Errorf("chi-square %v", result.ChiSquare)
|
||||||
|
}
|
||||||
|
if math.IsNaN(result.PValue) || result.PValue > 1 || result.PValue < 0 {
|
||||||
|
t.Errorf("p-value %v", result.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// With no events anywhere the covariance matrix is singular and there is
|
||||||
|
// nothing to test, which must report no difference rather than a divide by zero.
|
||||||
|
func TestLogRank_NoEventsAnywhere(t *testing.T) {
|
||||||
|
result := logRank([]string{"a", "b"}, [][]observation{
|
||||||
|
{{40, false}, {40, false}},
|
||||||
|
{{40, false}, {40, false}, {40, false}},
|
||||||
|
})
|
||||||
|
if result.ChiSquare != 0 || result.PValue != 1 {
|
||||||
|
t.Errorf("chi-square %v p-value %v, want 0 and 1", result.ChiSquare, result.PValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in new issue
Block a user