827 lines
25 KiB
Go
827 lines
25 KiB
Go
// 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
|
|||
|
|
}
|