feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+826
@@ -0,0 +1,826 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package stats
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The linear mixed model: a fixed-effect design shared by every
|
||||
// observation plus a random-effect design whose coefficients vary by
|
||||
// group, y = X·β + Z·b + ε with b_g ~ N(0, Σ) per group and
|
||||
// ε ~ N(0, σ²I). The variance components (Σ and σ²) are estimated by
|
||||
// residual maximum likelihood, the REML criterion that compares only
|
||||
// the error contrasts and so does not let the fixed effects soak up
|
||||
// degrees of freedom the variance estimate needs; the fixed effects
|
||||
// follow by general least squares at the fitted components.
|
||||
//
|
||||
// The maximisation is the expectation-maximisation sweep over the
|
||||
// random-effect posterior: every update is a conditional expectation,
|
||||
// so the sweep climbs the log likelihood monotonically until the
|
||||
// tolerance stops it, and the whole fit is deterministic for a given
|
||||
// input. No generator enters: the starting point comes from the
|
||||
// data's own ordinary least squares, and the sweep walks from there.
|
||||
|
||||
// mixedMaxIterations caps the EM sweeps; a fit that has not settled by
|
||||
// then is reported with Converged false rather than pretended to have
|
||||
// converged.
|
||||
const mixedMaxIterations = 500
|
||||
|
||||
// mixedTolerance is the convergence tolerance on the REML log
|
||||
// likelihood: the sweeps stop once it moves by less than
|
||||
// mixedTolerance scaled by 1 + |log likelihood|, past which no
|
||||
// parameter the M step can produce moves the objective materially.
|
||||
const mixedTolerance = 1e-10
|
||||
|
||||
// LinearMixedModelResult carries the fit of a linear mixed model.
|
||||
type LinearMixedModelResult struct {
|
||||
// Coefficients are the fitted fixed effects β̂, one per column of
|
||||
// the fixed design, and StandardErrors their estimated standard
|
||||
// deviations, the square roots of the diagonal of the GLS
|
||||
// covariance (XᵀV⁻¹X)⁻¹ at the fitted variance components.
|
||||
Coefficients []float64
|
||||
StandardErrors []float64
|
||||
// RandomEffects holds one coefficient vector per group, in the
|
||||
// order GroupLabels names, each of the length of a row of the
|
||||
// random design: the posterior means b̂_g at the fitted components.
|
||||
RandomEffects [][]float64
|
||||
// GroupLabels are the distinct group labels in the order the fit
|
||||
// met them, the order RandomEffects keeps.
|
||||
GroupLabels []int
|
||||
// RandomCovariance is the fitted between-group covariance Σ̂ of the
|
||||
// random effects, row-major, and ResidualVariance the fitted σ̂².
|
||||
RandomCovariance []float64
|
||||
ResidualVariance float64
|
||||
// LogLikelihood is the maximised REML log likelihood, the
|
||||
// criterion the sweep climbed.
|
||||
LogLikelihood float64
|
||||
// Fitted and Residuals align with the rows of the design: the
|
||||
// conditional fit X·β̂ + Z·b̂ per row and its remainder.
|
||||
Fitted []float64
|
||||
Residuals []float64
|
||||
// Iterations counts the EM sweeps taken; Converged reports whether
|
||||
// the log likelihood settled under mixedTolerance before the
|
||||
// budget ran out.
|
||||
Iterations int
|
||||
Converged bool
|
||||
}
|
||||
|
||||
// mixedGroup holds one group's fixed pieces, read once before the
|
||||
// sweep: the row indices, both designs restricted to the group, the
|
||||
// response restricted to it, and the random design's own scatter
|
||||
// ZᵀZ, which no iteration moves.
|
||||
type mixedGroup struct {
|
||||
label int
|
||||
rows []int
|
||||
yg []float64
|
||||
xg []float64
|
||||
zg []float64
|
||||
ztz []float64
|
||||
xtz []float64
|
||||
ng int
|
||||
}
|
||||
|
||||
// LinearMixedModel fits the linear mixed model over the response y,
|
||||
// the fixed design x (n rows, p columns, the intercept included by
|
||||
// the caller as a constant column when one is wanted), the random
|
||||
// design z (n rows, q columns) and the group label of every row. The
|
||||
// labels need no particular order or contiguity: the fit meets them
|
||||
// in row order and reports the distinct ones in GroupLabels in the
|
||||
// order it met them. The random effects share one unstructured q×q
|
||||
// covariance Σ across the groups.
|
||||
//
|
||||
// The fit runs expectation maximisation over the random-effect
|
||||
// posterior to the documented tolerance, within mixedMaxIterations
|
||||
// sweeps, and reports the REML log likelihood of the fitted state. A
|
||||
// group whose marginal covariance fails to factor under the fitted
|
||||
// components, or a fixed design that is singular under them, is an
|
||||
// error naming the group or the design.
|
||||
//
|
||||
// Refuses complex input and any non-finite entry, a response shorter
|
||||
// than two rows, a design of fewer than one column, a row count
|
||||
// mismatch, more fixed coefficients than observations, and a
|
||||
// negative or missing group label.
|
||||
func LinearMixedModel(y, x, z *core.Array, groups []int) (*LinearMixedModelResult, error) {
|
||||
const name = "LinearMixedModel"
|
||||
if y.NDim() != 1 {
|
||||
return nil, base.Errf("%s: the response must be rank 1, got shape %s", name, base.ShapeText(y.Shape()))
|
||||
}
|
||||
if x.NDim() != 2 || z.NDim() != 2 {
|
||||
return nil, base.Errf("%s: both designs must be rank 2", name)
|
||||
}
|
||||
if y.Dtype() == core.Complex || x.Dtype() == core.Complex || z.Dtype() == core.Complex {
|
||||
return nil, base.Errf("%s: complex inputs are not supported", name)
|
||||
}
|
||||
n := y.Len()
|
||||
p := x.Shape()[1]
|
||||
q := z.Shape()[1]
|
||||
if x.Shape()[0] != n || z.Shape()[0] != n {
|
||||
return nil, base.Errf("%s: the response holds %d rows, the designs %d and %d", name, n, x.Shape()[0], z.Shape()[0])
|
||||
}
|
||||
if n < 2 {
|
||||
return nil, base.Errf("%s: at least two observations are needed, got %d", name, n)
|
||||
}
|
||||
if p < 1 || q < 1 {
|
||||
return nil, base.Errf("%s: both designs need at least one column, got %d and %d", name, p, q)
|
||||
}
|
||||
if n <= p {
|
||||
return nil, base.Errf("%s: %d observations cannot carry %d fixed coefficients", name, n, p)
|
||||
}
|
||||
if len(groups) != n {
|
||||
return nil, base.Errf("%s: %d group labels for %d observations", name, len(groups), n)
|
||||
}
|
||||
if err := checkFinite(name, "the response", y); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := checkFinite(name, "the fixed design", x); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := checkFinite(name, "the random design", z); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// The response and the designs are read once into plain slices:
|
||||
// every sweep below indexes them flat.
|
||||
yVals := make([]float64, n)
|
||||
if fs := rawFloats(y); fs != nil {
|
||||
copy(yVals, fs)
|
||||
} else {
|
||||
for i := range n {
|
||||
yVals[i] = y.FloatAt(i)
|
||||
}
|
||||
}
|
||||
xVals := make([]float64, n*p)
|
||||
if fs := rawFloats(x); fs != nil {
|
||||
copy(xVals, fs[:n*p])
|
||||
} else {
|
||||
for i := range xVals {
|
||||
xVals[i] = x.FloatAt(i)
|
||||
}
|
||||
}
|
||||
zVals := make([]float64, n*q)
|
||||
if fs := rawFloats(z); fs != nil {
|
||||
copy(zVals, fs[:n*q])
|
||||
} else {
|
||||
for i := range zVals {
|
||||
zVals[i] = z.FloatAt(i)
|
||||
}
|
||||
}
|
||||
// The groups are canonicalised by first appearance: the labels
|
||||
// stay the caller's own, the fit only needs them distinct and
|
||||
// stable.
|
||||
labelIndex := make(map[int]int)
|
||||
labels := make([]int, 0, 8)
|
||||
member := make([][]int, 0, 8)
|
||||
for i, label := range groups {
|
||||
if label < 0 {
|
||||
return nil, base.Errf("%s: group label %d is negative", name, label)
|
||||
}
|
||||
gi, ok := labelIndex[label]
|
||||
if !ok {
|
||||
gi = len(labels)
|
||||
labelIndex[label] = gi
|
||||
labels = append(labels, label)
|
||||
member = append(member, nil)
|
||||
}
|
||||
member[gi] = append(member[gi], i)
|
||||
}
|
||||
gs := make([]mixedGroup, len(labels))
|
||||
for gi, rows := range member {
|
||||
g := &gs[gi]
|
||||
g.label = labels[gi]
|
||||
g.rows = rows
|
||||
g.ng = len(rows)
|
||||
g.yg = make([]float64, g.ng)
|
||||
g.xg = make([]float64, g.ng*p)
|
||||
g.zg = make([]float64, g.ng*q)
|
||||
g.ztz = make([]float64, q*q)
|
||||
g.xtz = make([]float64, p*q)
|
||||
for li, row := range rows {
|
||||
g.yg[li] = yVals[row]
|
||||
copy(g.xg[li*p:(li+1)*p], xVals[row*p:(row+1)*p])
|
||||
copy(g.zg[li*q:(li+1)*q], zVals[row*q:(row+1)*q])
|
||||
}
|
||||
for a := range q {
|
||||
for b := range q {
|
||||
s := 0.0
|
||||
for li := range g.ng {
|
||||
s += g.zg[li*q+a] * g.zg[li*q+b]
|
||||
}
|
||||
g.ztz[a*q+b] = s
|
||||
}
|
||||
}
|
||||
for j := range p {
|
||||
for b := range q {
|
||||
s := 0.0
|
||||
for li := range g.ng {
|
||||
s += g.xg[li*p+j] * g.zg[li*q+b]
|
||||
}
|
||||
g.xtz[j*q+b] = s
|
||||
}
|
||||
}
|
||||
}
|
||||
// The starting point: ordinary least squares on the fixed design,
|
||||
// the residual variance around it, and a between-group scatter of
|
||||
// its residuals floored above zero so the first sweep can see the
|
||||
// random effects at all.
|
||||
beta, err := mixedOLSStart(name, xVals, yVals, n, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
residual := make([]float64, n)
|
||||
rss := 0.0
|
||||
totalSquare := 0.0
|
||||
for i := range n {
|
||||
fit := 0.0
|
||||
for j := range p {
|
||||
fit += xVals[i*p+j] * beta[j]
|
||||
}
|
||||
residual[i] = yVals[i] - fit
|
||||
rss += residual[i] * residual[i]
|
||||
totalSquare += yVals[i] * yVals[i]
|
||||
}
|
||||
scale := 1 + totalSquare/float64(n)
|
||||
sigma2 := math.Max(rss/float64(n), 1e-12*scale)
|
||||
between := 0.0
|
||||
if len(gs) > 1 {
|
||||
overall := 0.0
|
||||
means := make([]float64, len(gs))
|
||||
for gi, g := range gs {
|
||||
s := 0.0
|
||||
for _, row := range g.rows {
|
||||
s += residual[row]
|
||||
}
|
||||
means[gi] = s / float64(g.ng)
|
||||
overall += means[gi]
|
||||
}
|
||||
overall /= float64(len(gs))
|
||||
for _, m := range means {
|
||||
d := m - overall
|
||||
between += d * d
|
||||
}
|
||||
between /= float64(len(gs) - 1)
|
||||
}
|
||||
comp := math.Max(between, 0.1*sigma2)
|
||||
sigma := make([]float64, q*q)
|
||||
for a := range q {
|
||||
sigma[a*q+a] = comp
|
||||
}
|
||||
// The sweep. The per-group buffers a pass touches are allocated
|
||||
// once here and refilled in place; the fixed-effect solve keeps
|
||||
// its own small p×p workspace, reallocated each pass.
|
||||
chol := make([][][]float64, len(gs))
|
||||
scratch := make([]mixedSweep, len(gs))
|
||||
for gi := range gs {
|
||||
ng := gs[gi].ng
|
||||
chol[gi] = make([][]float64, ng)
|
||||
for i := range ng {
|
||||
chol[gi][i] = make([]float64, ng)
|
||||
}
|
||||
scratch[gi] = newMixedSweep(ng, p, q)
|
||||
}
|
||||
// The fixed scatter XᵀX, which no iteration moves: the REML M
|
||||
// step's expectation of the squared residual carries its trace
|
||||
// against the fixed effects' posterior covariance.
|
||||
xtx := make([]float64, p*p)
|
||||
for i := range n {
|
||||
for j := range p {
|
||||
xj := xVals[i*p+j]
|
||||
for k := range p {
|
||||
xtx[j*p+k] += xj * xVals[i*p+k]
|
||||
}
|
||||
}
|
||||
}
|
||||
xtvix := make([]float64, p*p)
|
||||
xtviy := make([]float64, p)
|
||||
sigmaNew := make([]float64, q*q)
|
||||
logLik := math.Inf(-1)
|
||||
converged := false
|
||||
iterations := mixedMaxIterations
|
||||
for iter := 1; iter <= mixedMaxIterations; iter++ {
|
||||
pieces, perr := mixedPass(gs, chol, scratch, xtx, beta, sigma, sigma2, xtvix, xtviy, sigmaNew)
|
||||
if perr != nil {
|
||||
return nil, base.Errf("%s: %w", name, perr)
|
||||
}
|
||||
next := -0.5 * (pieces.logDetV + pieces.quad + pieces.logDetX + float64(n-p)*math.Log(2*math.Pi))
|
||||
if math.IsNaN(next) {
|
||||
return nil, base.Errf("%s: the fit left the finite domain at iteration %d", name, iter)
|
||||
}
|
||||
move := next - logLik
|
||||
logLik = next
|
||||
if iter > 1 && math.Abs(move) <= mixedTolerance*(1+math.Abs(logLik)) {
|
||||
converged = true
|
||||
iterations = iter
|
||||
break
|
||||
}
|
||||
// The M step: the averaged posterior scatter, the expected
|
||||
// residual variance with its posterior-trace correction, and
|
||||
// the GLS fixed effects at this pass's components.
|
||||
copy(sigma, sigmaNew)
|
||||
for a := range q {
|
||||
for b := a + 1; b < q; b++ {
|
||||
sigma[b*q+a] = sigma[a*q+b]
|
||||
}
|
||||
}
|
||||
sigma2 = math.Max((pieces.sse+pieces.trace)/float64(n), 1e-12*scale)
|
||||
// The divergence fuse: some designs give the REML criterion no
|
||||
// interior optimum, the classical case being a random design
|
||||
// that spans the fixed one under an unstructured covariance,
|
||||
// where the surface climbs a ridge of singular Σ without a
|
||||
// summit. The fit refuses with the condition named instead of
|
||||
// publishing the climb.
|
||||
if sigma2 > 1e12*scale {
|
||||
return nil, base.Errf("%s: the variance components diverged past the data's scale; this design gives the REML criterion no interior optimum", name)
|
||||
}
|
||||
for _, v := range sigma {
|
||||
if math.Abs(v) > 1e12*scale {
|
||||
return nil, base.Errf("%s: the variance components diverged past the data's scale; this design gives the REML criterion no interior optimum", name)
|
||||
}
|
||||
}
|
||||
if err := lltSolveSquare(name, xtvix, xtviy, p); err != nil {
|
||||
return nil, base.Errf("%s: the fixed design is singular under the fitted covariance (%w)", name, err)
|
||||
}
|
||||
copy(beta, xtviy)
|
||||
}
|
||||
// The reporting pass at the fitted parameters: the same sweep once
|
||||
// more at fixed state, then the quantities the caller holds. The
|
||||
// sweep is deterministic, so its totals are the fitted state's
|
||||
// own.
|
||||
pieces, perr := mixedPass(gs, chol, scratch, xtx, beta, sigma, sigma2, xtvix, xtviy, sigmaNew)
|
||||
if perr != nil {
|
||||
return nil, base.Errf("%s: %w", name, perr)
|
||||
}
|
||||
logLik = -0.5*(pieces.logDetV+pieces.quad+pieces.logDetX) - 0.5*float64(n-p)*math.Log(2*math.Pi)
|
||||
// The standard errors: the square roots of the diagonal of the GLS
|
||||
// covariance, from the inverse of the same matrix the log
|
||||
// determinant came from.
|
||||
diag, serr := lltInverseDiagonal(name, xtvix, p)
|
||||
if serr != nil {
|
||||
return nil, base.Errf("%s: %w", name, serr)
|
||||
}
|
||||
se := make([]float64, p)
|
||||
for i := range p {
|
||||
se[i] = math.Sqrt(diag[i])
|
||||
}
|
||||
fitted := make([]float64, n)
|
||||
resids := make([]float64, n)
|
||||
for gi, g := range gs {
|
||||
bhat := scratch[gi].bhat
|
||||
for li, row := range g.rows {
|
||||
fit := 0.0
|
||||
for j := range p {
|
||||
fit += xVals[row*p+j] * beta[j]
|
||||
}
|
||||
for a := range q {
|
||||
fit += g.zg[li*q+a] * bhat[a]
|
||||
}
|
||||
fitted[row] = fit
|
||||
resids[row] = yVals[row] - fit
|
||||
}
|
||||
}
|
||||
effects := make([][]float64, len(gs))
|
||||
for gi := range gs {
|
||||
effects[gi] = append([]float64(nil), scratch[gi].bhat...)
|
||||
}
|
||||
return &LinearMixedModelResult{
|
||||
Coefficients: append([]float64(nil), beta...),
|
||||
StandardErrors: se,
|
||||
RandomEffects: effects,
|
||||
GroupLabels: labels,
|
||||
RandomCovariance: append([]float64(nil), sigma...),
|
||||
ResidualVariance: sigma2,
|
||||
LogLikelihood: logLik,
|
||||
Fitted: fitted,
|
||||
Residuals: resids,
|
||||
Iterations: iterations,
|
||||
Converged: converged,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// mixedSweep holds one group's reusable pass buffers: the marginal
|
||||
// covariance and its factor's solves against the response and both
|
||||
// designs, the conditional residual, and the posterior mean and
|
||||
// scatter of the group's random effects.
|
||||
type mixedSweep struct {
|
||||
u []float64 // V⁻¹ r
|
||||
vy []float64 // V⁻¹ y
|
||||
vz []float64 // V⁻¹ Z, ng×q
|
||||
vx []float64 // V⁻¹ X, ng×p
|
||||
slope []float64 // −Σ Zᵀ (V⁻¹ X), the posterior mean's slope in β, q×p
|
||||
bhat []float64 // Σ Zᵀ u
|
||||
post []float64 // P + b̂b̂ᵀ, the conditional scatter
|
||||
v []float64 // the marginal covariance, ng×ng
|
||||
resid []float64 // y − Xβ per row
|
||||
mzx []float64 // Σ Zᵀ (V⁻¹ Z), q×q
|
||||
sc []float64 // slope·C, refilled by the REML correction, q×p
|
||||
bq []float64 // sc·slopeᵀ, the correction's scatter, q×q
|
||||
colvec []float64 // one design column, refilled per solve
|
||||
}
|
||||
|
||||
// newMixedSweep allocates one group's buffers.
|
||||
func newMixedSweep(ng, p, q int) mixedSweep {
|
||||
return mixedSweep{
|
||||
u: make([]float64, ng),
|
||||
vy: make([]float64, ng),
|
||||
vz: make([]float64, ng*q),
|
||||
vx: make([]float64, ng*p),
|
||||
slope: make([]float64, q*p),
|
||||
bhat: make([]float64, q),
|
||||
post: make([]float64, q*q),
|
||||
v: make([]float64, ng*ng),
|
||||
resid: make([]float64, ng),
|
||||
mzx: make([]float64, q*q),
|
||||
sc: make([]float64, q*p),
|
||||
bq: make([]float64, q*q),
|
||||
colvec: make([]float64, ng),
|
||||
}
|
||||
}
|
||||
|
||||
// mixedPassTotals gathers what one pass accumulates across the
|
||||
// groups.
|
||||
type mixedPassTotals struct {
|
||||
logDetV float64 // Σ log|V_g|, the marginal covariances' own logarithms
|
||||
quad float64 // Σ r_gᵀ V_g⁻¹ r_g
|
||||
sse float64 // Σ ‖r_g − Z_g b̂_g‖²
|
||||
trace float64 // Σ tr(Z_gᵀZ_g·P_g)
|
||||
logDetX float64 // log|XᵀV⁻¹X|
|
||||
}
|
||||
|
||||
// mixedPass runs one expectation pass over the groups at the given
|
||||
// parameters: it factors every group's marginal covariance, solves it
|
||||
// against the response and both designs, forms the posterior means
|
||||
// and scatters, and accumulates the totals the log likelihood and the
|
||||
// M step read. The M step's expectations are taken under the flat
|
||||
// prior on the fixed effects, so the fixed effects' own posterior
|
||||
// covariance (XᵀV⁻¹X)⁻¹ joins both updates through the posterior
|
||||
// mean's slope in β; that is the step that makes the fixed point the
|
||||
// REML optimum rather than the ML one. The accumulators xtvix, xtviy
|
||||
// and sigmaNew are overwritten; the posterior means stay in the
|
||||
// group's scratch for the reporting pass to collect.
|
||||
func mixedPass(gs []mixedGroup, chol [][][]float64, scratch []mixedSweep, xtx []float64,
|
||||
beta []float64, sigma []float64, sigma2 float64, xtvix, xtviy, sigmaNew []float64) (*mixedPassTotals, error) {
|
||||
const name = "LinearMixedModel"
|
||||
p := len(beta)
|
||||
q := len(scratch[0].bhat)
|
||||
clear(xtvix)
|
||||
clear(xtviy)
|
||||
clear(sigmaNew)
|
||||
totals := &mixedPassTotals{}
|
||||
for gi := range gs {
|
||||
g := &gs[gi]
|
||||
s := &scratch[gi]
|
||||
ng := g.ng
|
||||
// The marginal covariance V = Z Σ Zᵀ + σ²I, both triangles
|
||||
// written so the factorisation's symmetry check sees a mirror
|
||||
// pair with identical bits.
|
||||
for i := range ng {
|
||||
for j := range i + 1 {
|
||||
total := 0.0
|
||||
for a := range q {
|
||||
za := g.zg[i*q+a]
|
||||
if za != 0 {
|
||||
for b := range q {
|
||||
total += za * sigma[a*q+b] * g.zg[j*q+b]
|
||||
}
|
||||
}
|
||||
}
|
||||
if i == j {
|
||||
total += sigma2
|
||||
}
|
||||
s.v[i*ng+j] = total
|
||||
s.v[j*ng+i] = total
|
||||
}
|
||||
}
|
||||
if err := mvnCholeskyFlatInto(name, s.v, chol[gi], ng); err != nil {
|
||||
return nil, base.Errf("group %d's marginal covariance failed to factor (%w)", g.label, err)
|
||||
}
|
||||
l := chol[gi]
|
||||
for i := range ng {
|
||||
// The factor's determinant answers |V_g| = |L|², so the
|
||||
// diagonal logs enter doubled: the criterion the sweep reads
|
||||
// below carries log|V_g| itself, not half of it.
|
||||
totals.logDetV += 2 * math.Log(l[i][i])
|
||||
}
|
||||
// The residual against the fixed part alone, then u = V⁻¹ r.
|
||||
for li := range ng {
|
||||
r := 0.0
|
||||
for j := range p {
|
||||
r += g.xg[li*p+j] * beta[j]
|
||||
}
|
||||
s.resid[li] = g.yg[li] - r
|
||||
}
|
||||
copy(s.u, s.resid)
|
||||
lltSolveInPlace(l, s.u)
|
||||
totals.quad += dot(s.resid, s.u)
|
||||
// V⁻¹ Z and V⁻¹ X, column by column through the same factor.
|
||||
for a := range q {
|
||||
for li := range ng {
|
||||
s.colvec[li] = g.zg[li*q+a]
|
||||
}
|
||||
lltSolveInPlace(l, s.colvec)
|
||||
copy(s.vz[a*ng:a*ng+ng], s.colvec)
|
||||
}
|
||||
for j := range p {
|
||||
for li := range ng {
|
||||
s.colvec[li] = g.xg[li*p+j]
|
||||
}
|
||||
lltSolveInPlace(l, s.colvec)
|
||||
copy(s.vx[j*ng:j*ng+ng], s.colvec)
|
||||
}
|
||||
// V⁻¹ y, free of charge out of the solves already run:
|
||||
// V⁻¹ y = V⁻¹ r + (V⁻¹ X)·β.
|
||||
for li := range ng {
|
||||
total := s.u[li]
|
||||
for j := range p {
|
||||
total += s.vx[j*ng+li] * beta[j]
|
||||
}
|
||||
s.vy[li] = total
|
||||
}
|
||||
// The posterior mean b̂ = Σ Zᵀ u.
|
||||
clear(s.bhat)
|
||||
for a := range q {
|
||||
total := 0.0
|
||||
for li := range ng {
|
||||
zt := 0.0
|
||||
for b := range q {
|
||||
zt += sigma[a*q+b] * g.zg[li*q+b]
|
||||
}
|
||||
total += zt * s.u[li]
|
||||
}
|
||||
s.bhat[a] = total
|
||||
}
|
||||
// mzx = Σ Zᵀ (V⁻¹ Z), then the posterior covariance
|
||||
// P = Σ − mzx·Σ on its upper triangle, joined by the mean's
|
||||
// outer product into the scatter the M step averages.
|
||||
clear(s.mzx)
|
||||
for a := range q {
|
||||
for b := range q {
|
||||
total := 0.0
|
||||
for c := range q {
|
||||
factor := sigma[a*q+c]
|
||||
if factor != 0 {
|
||||
for li := range ng {
|
||||
total += factor * g.zg[li*q+c] * s.vz[li*q+b]
|
||||
}
|
||||
}
|
||||
}
|
||||
s.mzx[a*q+b] = total
|
||||
}
|
||||
}
|
||||
clear(s.post)
|
||||
for a := range q {
|
||||
for b := a; b < q; b++ {
|
||||
pv := sigma[a*q+b]
|
||||
for c := range q {
|
||||
pv -= s.mzx[a*q+c] * sigma[c*q+b]
|
||||
}
|
||||
s.post[a*q+b] = pv + s.bhat[a]*s.bhat[b]
|
||||
}
|
||||
}
|
||||
for a := range q {
|
||||
for b := a; b < q; b++ {
|
||||
sigmaNew[a*q+b] += s.post[a*q+b] / float64(len(gs))
|
||||
}
|
||||
}
|
||||
// The conditional residual and its posterior-trace correction.
|
||||
for li := range ng {
|
||||
r := s.resid[li]
|
||||
for a := range q {
|
||||
r -= g.zg[li*q+a] * s.bhat[a]
|
||||
}
|
||||
s.resid[li] = r
|
||||
totals.sse += r * r
|
||||
}
|
||||
for a := range q {
|
||||
for b := a; b < q; b++ {
|
||||
pv := sigma[a*q+b]
|
||||
for c := range q {
|
||||
pv -= s.mzx[a*q+c] * sigma[c*q+b]
|
||||
}
|
||||
totals.trace += g.ztz[a*q+b] * pv
|
||||
if b != a {
|
||||
totals.trace += g.ztz[b*q+a] * pv
|
||||
}
|
||||
}
|
||||
}
|
||||
// XᵀV⁻¹X and XᵀV⁻¹y accumulate across the groups. The
|
||||
// information is symmetric by construction but its two rounding
|
||||
// paths differ, so only the upper triangle is accumulated and
|
||||
// mirrored: the factorisation's relative symmetry check would
|
||||
// otherwise condemn a matrix whose one side rounded to an exact
|
||||
// zero.
|
||||
for j := range p {
|
||||
for k := j; k < p; k++ {
|
||||
total := 0.0
|
||||
for li := range ng {
|
||||
total += g.xg[li*p+k] * s.vx[j*ng+li]
|
||||
}
|
||||
xtvix[j*p+k] += total
|
||||
if k != j {
|
||||
xtvix[k*p+j] += total
|
||||
}
|
||||
}
|
||||
total := 0.0
|
||||
for li := range ng {
|
||||
total += g.xg[li*p+j] * s.vy[li]
|
||||
}
|
||||
xtviy[j] += total
|
||||
}
|
||||
// The posterior mean's slope in β: E[b_g|β, y] moves as
|
||||
// b̂ − (Σ ZᵀV⁻¹X)(β − β̂), and the slope carries the fixed
|
||||
// effects' uncertainty into the M step's expectations.
|
||||
for a := range q {
|
||||
for j := range p {
|
||||
total := 0.0
|
||||
for b := range q {
|
||||
zt := sigma[a*q+b]
|
||||
if zt != 0 {
|
||||
for li := range ng {
|
||||
total += zt * g.zg[li*q+b] * s.vx[j*ng+li]
|
||||
}
|
||||
}
|
||||
}
|
||||
s.slope[a*p+j] = -total
|
||||
}
|
||||
}
|
||||
}
|
||||
// The fixed information through its own Cholesky: the log
|
||||
// determinant the REML criterion reads, and the posterior
|
||||
// covariance C = (XᵀV⁻¹X)⁻¹ the corrections below walk through.
|
||||
// A non-positive pivot means the design has collapsed under the
|
||||
// fitted covariance, which is an error, not a fit.
|
||||
cholX := make([][]float64, p)
|
||||
for i := range p {
|
||||
cholX[i] = make([]float64, p)
|
||||
}
|
||||
if err := mvnCholeskyFlatInto(name, xtvix, cholX, p); err != nil {
|
||||
return nil, base.Errf("the fixed design is singular under the fitted covariance (%w)", err)
|
||||
}
|
||||
logDetX := 0.0
|
||||
for i := range p {
|
||||
logDetX += math.Log(cholX[i][i])
|
||||
}
|
||||
totals.logDetX = 2 * logDetX
|
||||
inverse := make([]float64, p*p)
|
||||
column := make([]float64, p)
|
||||
for j := range p {
|
||||
clear(column)
|
||||
column[j] = 1
|
||||
lltSolveInPlace(cholX, column)
|
||||
for i := range p {
|
||||
inverse[i*p+j] = column[i]
|
||||
}
|
||||
}
|
||||
// The flat-prior corrections. The expected squared residual gains
|
||||
// the trace of XᵀX against C and, per group, twice the cross term
|
||||
// of XᵀZ against C·slopeᵀ; the expected random-effect scatter
|
||||
// gains slope·C·slopeᵀ. With no random effect the two terms leave
|
||||
// σ² = RSS/(n−p), the REML answer, as the fixed point.
|
||||
traceXX := 0.0
|
||||
for a := range p {
|
||||
for b := range p {
|
||||
traceXX += xtx[a*p+b] * inverse[a*p+b]
|
||||
}
|
||||
}
|
||||
crossTotal := 0.0
|
||||
for gi := range gs {
|
||||
g := &gs[gi]
|
||||
s := &scratch[gi]
|
||||
for a := range q {
|
||||
for c := range p {
|
||||
total := 0.0
|
||||
for j := range p {
|
||||
total += s.slope[a*p+j] * inverse[j*p+c]
|
||||
}
|
||||
s.sc[a*p+c] = total
|
||||
}
|
||||
}
|
||||
clear(s.bq)
|
||||
for a := range q {
|
||||
for b := a; b < q; b++ {
|
||||
total := 0.0
|
||||
for c := range p {
|
||||
total += s.sc[a*p+c] * s.slope[b*p+c]
|
||||
}
|
||||
s.bq[a*q+b] = total
|
||||
}
|
||||
}
|
||||
for a := range q {
|
||||
for b := a; b < q; b++ {
|
||||
sigmaNew[a*q+b] += s.bq[a*q+b] / float64(len(gs))
|
||||
totals.trace += g.ztz[a*q+b] * s.bq[a*q+b]
|
||||
if b != a {
|
||||
totals.trace += g.ztz[b*q+a] * s.bq[a*q+b]
|
||||
}
|
||||
}
|
||||
}
|
||||
for j := range p {
|
||||
for b := range q {
|
||||
crossTotal += g.xtz[j*q+b] * s.sc[b*p+j]
|
||||
}
|
||||
}
|
||||
}
|
||||
totals.sse += traceXX + 2*crossTotal
|
||||
return totals, nil
|
||||
}
|
||||
|
||||
// dot returns the plain dot product of two equal-length slices.
|
||||
func dot(a, b []float64) float64 {
|
||||
total := 0.0
|
||||
for i, v := range a {
|
||||
total += v * b[i]
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// lltSolveInPlace solves L·Lᵀ x = x in place: the forward sweep
|
||||
// leaves L's solution in x, the backward sweep reads the entries the
|
||||
// descending order has already updated.
|
||||
func lltSolveInPlace(l [][]float64, x []float64) {
|
||||
n := len(x)
|
||||
for i := range n {
|
||||
total := x[i]
|
||||
for j := range i {
|
||||
total -= l[i][j] * x[j]
|
||||
}
|
||||
x[i] = total / l[i][i]
|
||||
}
|
||||
for i := n - 1; i >= 0; i-- {
|
||||
total := x[i]
|
||||
for j := i + 1; j < n; j++ {
|
||||
total -= l[j][i] * x[j]
|
||||
}
|
||||
x[i] = total / l[i][i]
|
||||
}
|
||||
}
|
||||
|
||||
// lltSolveSquare solves the symmetric positive definite system a·x =
|
||||
// b in place, factoring a through the house Cholesky and leaving b
|
||||
// holding x.
|
||||
func lltSolveSquare(name string, a, b []float64, d int) error {
|
||||
l := make([][]float64, d)
|
||||
for i := range d {
|
||||
l[i] = make([]float64, d)
|
||||
}
|
||||
if err := mvnCholeskyFlatInto(name, a, l, d); err != nil {
|
||||
return err
|
||||
}
|
||||
lltSolveInPlace(l, b)
|
||||
return nil
|
||||
}
|
||||
|
||||
// lltInverseDiagonal returns the diagonal of the inverse of the
|
||||
// symmetric positive definite a, one solve against each unit vector.
|
||||
func lltInverseDiagonal(name string, a []float64, d int) ([]float64, error) {
|
||||
l := make([][]float64, d)
|
||||
for i := range d {
|
||||
l[i] = make([]float64, d)
|
||||
}
|
||||
if err := mvnCholeskyFlatInto(name, a, l, d); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
diag := make([]float64, d)
|
||||
column := make([]float64, d)
|
||||
for i := range d {
|
||||
clear(column)
|
||||
column[i] = 1
|
||||
lltSolveInPlace(l, column)
|
||||
diag[i] = column[i]
|
||||
}
|
||||
return diag, nil
|
||||
}
|
||||
|
||||
// mixedOLSStart returns the ordinary least squares coefficients of
|
||||
// the fixed design, the fixed effects' starting point.
|
||||
func mixedOLSStart(name string, xVals, yVals []float64, n, p int) ([]float64, error) {
|
||||
xtx := make([]float64, p*p)
|
||||
xty := make([]float64, p)
|
||||
for i := range n {
|
||||
for j := range p {
|
||||
xj := xVals[i*p+j]
|
||||
xty[j] += xj * yVals[i]
|
||||
for k := range p {
|
||||
xtx[j*p+k] += xj * xVals[i*p+k]
|
||||
}
|
||||
}
|
||||
}
|
||||
rhs := make([][]float64, 1)
|
||||
rhs[0] = xty
|
||||
solved, err := base.SolveSystem(name, toRows(xtx, p), rhs)
|
||||
if err != nil {
|
||||
return nil, base.Errf("%s: the fixed design is singular, the fixed effects cannot be initialised (%w)", name, err)
|
||||
}
|
||||
return append([]float64(nil), solved[0]...), nil
|
||||
}
|
||||
|
||||
// toRows views a flat row-major matrix as rows for the shared solve.
|
||||
func toRows(flat []float64, d int) [][]float64 {
|
||||
rows := make([][]float64, d)
|
||||
for i := range d {
|
||||
rows[i] = flat[i*d : i*d+d]
|
||||
}
|
||||
return rows
|
||||
}
|
||||
Reference in New Issue
Block a user