From 019d608f659f950a311a0c186b23661073274504 Mon Sep 17 00:00:00 2001 From: PJ Date: Wed, 12 Aug 2026 23:03:03 +0530 Subject: [PATCH] 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 --- cmd/internal-tools/analyze/analysis.go | 192 +++++++++++ cmd/internal-tools/analyze/analysis_test.go | 218 +++++++++++++ cmd/internal-tools/analyze/distribution.go | 83 +++++ .../analyze/distribution_test.go | 74 +++++ cmd/internal-tools/analyze/end_to_end_test.go | 246 ++++++++++++++ cmd/internal-tools/analyze/fixtures_test.go | 55 ++++ cmd/internal-tools/analyze/holm.go | 36 +++ cmd/internal-tools/analyze/holm_test.go | 54 ++++ cmd/internal-tools/analyze/load.go | 242 ++++++++++++++ cmd/internal-tools/analyze/load_test.go | 163 ++++++++++ cmd/internal-tools/analyze/main.go | 102 ++++++ cmd/internal-tools/analyze/ranksum.go | 220 +++++++++++++ cmd/internal-tools/analyze/ranksum_test.go | 211 ++++++++++++ cmd/internal-tools/analyze/report.go | 154 +++++++++ cmd/internal-tools/analyze/survival.go | 241 ++++++++++++++ cmd/internal-tools/analyze/survival_test.go | 304 ++++++++++++++++++ 16 files changed, 2595 insertions(+) create mode 100644 cmd/internal-tools/analyze/analysis.go create mode 100644 cmd/internal-tools/analyze/analysis_test.go create mode 100644 cmd/internal-tools/analyze/distribution.go create mode 100644 cmd/internal-tools/analyze/distribution_test.go create mode 100644 cmd/internal-tools/analyze/end_to_end_test.go create mode 100644 cmd/internal-tools/analyze/fixtures_test.go create mode 100644 cmd/internal-tools/analyze/holm.go create mode 100644 cmd/internal-tools/analyze/holm_test.go create mode 100644 cmd/internal-tools/analyze/load.go create mode 100644 cmd/internal-tools/analyze/load_test.go create mode 100644 cmd/internal-tools/analyze/main.go create mode 100644 cmd/internal-tools/analyze/ranksum.go create mode 100644 cmd/internal-tools/analyze/ranksum_test.go create mode 100644 cmd/internal-tools/analyze/report.go create mode 100644 cmd/internal-tools/analyze/survival.go create mode 100644 cmd/internal-tools/analyze/survival_test.go diff --git a/cmd/internal-tools/analyze/analysis.go b/cmd/internal-tools/analyze/analysis.go new file mode 100644 index 0000000..35b6cc3 --- /dev/null +++ b/cmd/internal-tools/analyze/analysis.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/analysis_test.go b/cmd/internal-tools/analyze/analysis_test.go new file mode 100644 index 0000000..5311db6 --- /dev/null +++ b/cmd/internal-tools/analyze/analysis_test.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/distribution.go b/cmd/internal-tools/analyze/distribution.go new file mode 100644 index 0000000..e010b3d --- /dev/null +++ b/cmd/internal-tools/analyze/distribution.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/distribution_test.go b/cmd/internal-tools/analyze/distribution_test.go new file mode 100644 index 0000000..d020c31 --- /dev/null +++ b/cmd/internal-tools/analyze/distribution_test.go @@ -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) + } + } +} diff --git a/cmd/internal-tools/analyze/end_to_end_test.go b/cmd/internal-tools/analyze/end_to_end_test.go new file mode 100644 index 0000000..1ce33c0 --- /dev/null +++ b/cmd/internal-tools/analyze/end_to_end_test.go @@ -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") + } +} diff --git a/cmd/internal-tools/analyze/fixtures_test.go b/cmd/internal-tools/analyze/fixtures_test.go new file mode 100644 index 0000000..6092899 --- /dev/null +++ b/cmd/internal-tools/analyze/fixtures_test.go @@ -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} +) diff --git a/cmd/internal-tools/analyze/holm.go b/cmd/internal-tools/analyze/holm.go new file mode 100644 index 0000000..2d07e54 --- /dev/null +++ b/cmd/internal-tools/analyze/holm.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/holm_test.go b/cmd/internal-tools/analyze/holm_test.go new file mode 100644 index 0000000..79dfa5b --- /dev/null +++ b/cmd/internal-tools/analyze/holm_test.go @@ -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) + } +} diff --git a/cmd/internal-tools/analyze/load.go b/cmd/internal-tools/analyze/load.go new file mode 100644 index 0000000..b622f67 --- /dev/null +++ b/cmd/internal-tools/analyze/load.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/load_test.go b/cmd/internal-tools/analyze/load_test.go new file mode 100644 index 0000000..ce8f1fa --- /dev/null +++ b/cmd/internal-tools/analyze/load_test.go @@ -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) + } +} diff --git a/cmd/internal-tools/analyze/main.go b/cmd/internal-tools/analyze/main.go new file mode 100644 index 0000000..96d980a --- /dev/null +++ b/cmd/internal-tools/analyze/main.go @@ -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 ] [ ...] + +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) + } +} diff --git a/cmd/internal-tools/analyze/ranksum.go b/cmd/internal-tools/analyze/ranksum.go new file mode 100644 index 0000000..118df15 --- /dev/null +++ b/cmd/internal-tools/analyze/ranksum.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/ranksum_test.go b/cmd/internal-tools/analyze/ranksum_test.go new file mode 100644 index 0000000..4a61dbc --- /dev/null +++ b/cmd/internal-tools/analyze/ranksum_test.go @@ -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) + } + } +} diff --git a/cmd/internal-tools/analyze/report.go b/cmd/internal-tools/analyze/report.go new file mode 100644 index 0000000..47054b8 --- /dev/null +++ b/cmd/internal-tools/analyze/report.go @@ -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) + } +} diff --git a/cmd/internal-tools/analyze/survival.go b/cmd/internal-tools/analyze/survival.go new file mode 100644 index 0000000..b107f2a --- /dev/null +++ b/cmd/internal-tools/analyze/survival.go @@ -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 +} diff --git a/cmd/internal-tools/analyze/survival_test.go b/cmd/internal-tools/analyze/survival_test.go new file mode 100644 index 0000000..78e06cf --- /dev/null +++ b/cmd/internal-tools/analyze/survival_test.go @@ -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) + } +}