regressio 1.1.0 → 1.1.2

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/dist/index.cjs CHANGED
@@ -36,7 +36,17 @@ var __export = (target, all) => {
36
36
  set: __exportSetter.bind(all, name)
37
37
  });
38
38
  };
39
- var __esm = (fn, res) => () => (fn && (res = fn(fn = 0)), res);
39
+ var __esm = (fn, res, err) => () => {
40
+ if (fn)
41
+ try {
42
+ res = fn(fn = 0);
43
+ } catch (e) {
44
+ err = [e];
45
+ }
46
+ if (err)
47
+ throw err[0];
48
+ return res;
49
+ };
40
50
 
41
51
  // pkg/regressio_wasm_bg.js
42
52
  function bootstrap_ols(x, y, n, p, fit_intercept, n_bootstrap, seed) {
@@ -262,35 +272,35 @@ function __wbg_set_wasm(val) {
262
272
  var cachedFloat64ArrayMemory0 = null, WASM_VECTOR_LEN = 0, wasm;
263
273
 
264
274
  // pkg/regressio_wasm_bg.wasm
265
- var regressio_wasm_bg_default = "./regressio_wasm_bg-zc1d7yka.wasm";
275
+ var regressio_wasm_bg_default = "./regressio_wasm_bg-5tapzraw.wasm";
266
276
  var init_regressio_wasm_bg = () => {};
267
277
 
268
278
  // src/wasm-init.ts
269
279
  var exports_wasm_init = {};
270
280
  __export(exports_wasm_init, {
271
- vif: () => vif,
272
- vector_dot: () => vector_dot,
273
- svd: () => svd,
274
- solve_triangular: () => solve_triangular,
275
- softmax_rows: () => softmax_rows,
276
- qr_decompose: () => qr_decompose,
277
- matrix_transpose: () => matrix_transpose,
278
- matrix_subtract: () => matrix_subtract,
279
- matrix_scale: () => matrix_scale,
280
- matrix_multiply: () => matrix_multiply,
281
- matrix_add: () => matrix_add,
282
- manhattan_distances: () => manhattan_distances,
283
- irls_logistic: () => irls_logistic,
284
- initWasm: () => initWasm,
285
- frobenius_norm: () => frobenius_norm,
286
- forward_substitution: () => forward_substitution,
287
- euclidean_distances: () => euclidean_distances,
288
- eigenvalues: () => eigenvalues,
289
- determinant: () => determinant,
290
- correlation_matrix: () => correlation_matrix,
291
- coordinate_descent: () => coordinate_descent,
281
+ bootstrap_ols: () => bootstrap_ols,
292
282
  cholesky: () => cholesky,
293
- bootstrap_ols: () => bootstrap_ols
283
+ coordinate_descent: () => coordinate_descent,
284
+ correlation_matrix: () => correlation_matrix,
285
+ determinant: () => determinant,
286
+ eigenvalues: () => eigenvalues,
287
+ euclidean_distances: () => euclidean_distances,
288
+ forward_substitution: () => forward_substitution,
289
+ frobenius_norm: () => frobenius_norm,
290
+ initWasm: () => initWasm,
291
+ irls_logistic: () => irls_logistic,
292
+ manhattan_distances: () => manhattan_distances,
293
+ matrix_add: () => matrix_add,
294
+ matrix_multiply: () => matrix_multiply,
295
+ matrix_scale: () => matrix_scale,
296
+ matrix_subtract: () => matrix_subtract,
297
+ matrix_transpose: () => matrix_transpose,
298
+ qr_decompose: () => qr_decompose,
299
+ softmax_rows: () => softmax_rows,
300
+ solve_triangular: () => solve_triangular,
301
+ svd: () => svd,
302
+ vector_dot: () => vector_dot,
303
+ vif: () => vif
294
304
  });
295
305
  async function loadWasmBytes() {
296
306
  const resolved = typeof regressio_wasm_bg_default === "string" ? regressio_wasm_bg_default : null;
@@ -321,42 +331,42 @@ var init_wasm_init = __esm(() => {
321
331
  // src/index.ts
322
332
  var exports_src = {};
323
333
  __export(exports_src, {
324
- vif: () => vif2,
325
- unstandardize: () => unstandardize,
326
- unnormalize: () => unnormalize,
327
- studentizedResiduals: () => studentizedResiduals,
328
- standardize: () => standardize,
329
- shapiroWilk: () => shapiroWilk,
330
- residualDiagnostics: () => residualDiagnostics,
331
- predictionInterval: () => predictionInterval,
332
- polynomialFeatures: () => polynomialFeatures,
333
- oneHotEncode: () => oneHotEncode,
334
- normalize: () => normalize,
335
- leverage: () => leverage,
336
- isWasmActive: () => isWasmActive,
337
- interactionFeatures: () => interactionFeatures,
338
- imputeMedian: () => imputeMedian,
339
- imputeMean: () => imputeMean,
340
- durbinWatson: () => durbinWatson,
341
- dropMissing: () => dropMissing,
342
- correlationMatrix: () => correlationMatrix,
343
- cooksDistance: () => cooksDistance,
344
- confidenceInterval: () => confidenceInterval,
345
- conditionNumber: () => conditionNumber,
346
- breuschPagan: () => breuschPagan,
347
- bootstrapCoefficients: () => bootstrapCoefficients,
348
- WeightedRegression: () => WeightedRegression,
349
- RobustRegression: () => RobustRegression,
350
- RidgeRegression: () => RidgeRegression,
351
- PolynomialRegression: () => PolynomialRegression,
352
- NeuralNetwork: () => NeuralNetwork,
353
- MulticlassLogisticRegression: () => MulticlassLogisticRegression,
354
- Matrix: () => Matrix,
355
- LogisticRegression: () => LogisticRegression,
356
- LinearRegression: () => LinearRegression,
357
- LassoRegression: () => LassoRegression,
334
+ ElasticNet: () => ElasticNet,
358
335
  KNearestNeighbors: () => KNearestNeighbors,
359
- ElasticNet: () => ElasticNet
336
+ LassoRegression: () => LassoRegression,
337
+ LinearRegression: () => LinearRegression,
338
+ LogisticRegression: () => LogisticRegression,
339
+ Matrix: () => Matrix,
340
+ MulticlassLogisticRegression: () => MulticlassLogisticRegression,
341
+ NeuralNetwork: () => NeuralNetwork,
342
+ PolynomialRegression: () => PolynomialRegression,
343
+ RidgeRegression: () => RidgeRegression,
344
+ RobustRegression: () => RobustRegression,
345
+ WeightedRegression: () => WeightedRegression,
346
+ bootstrapCoefficients: () => bootstrapCoefficients,
347
+ breuschPagan: () => breuschPagan,
348
+ conditionNumber: () => conditionNumber,
349
+ confidenceInterval: () => confidenceInterval,
350
+ cooksDistance: () => cooksDistance,
351
+ correlationMatrix: () => correlationMatrix,
352
+ dropMissing: () => dropMissing,
353
+ durbinWatson: () => durbinWatson,
354
+ imputeMean: () => imputeMean,
355
+ imputeMedian: () => imputeMedian,
356
+ interactionFeatures: () => interactionFeatures,
357
+ isWasmActive: () => isWasmActive,
358
+ leverage: () => leverage,
359
+ normalize: () => normalize,
360
+ oneHotEncode: () => oneHotEncode,
361
+ polynomialFeatures: () => polynomialFeatures,
362
+ predictionInterval: () => predictionInterval,
363
+ residualDiagnostics: () => residualDiagnostics,
364
+ shapiroWilk: () => shapiroWilk,
365
+ standardize: () => standardize,
366
+ studentizedResiduals: () => studentizedResiduals,
367
+ unnormalize: () => unnormalize,
368
+ unstandardize: () => unstandardize,
369
+ vif: () => vif2
360
370
  });
361
371
  module.exports = __toCommonJS(exports_src);
362
372
 
@@ -1196,12 +1206,55 @@ function chi2TestPValue(stat, df) {
1196
1206
  return 1 - chi2CDFExact(stat, df);
1197
1207
  }
1198
1208
 
1209
+ // src/models/input-validation.ts
1210
+ function normalizeModelInput(X) {
1211
+ if (X.length === 0)
1212
+ throw new Error("Input data cannot be empty");
1213
+ const matrix = typeof X[0] === "number" ? X.map((value) => [value]) : X;
1214
+ const columns = matrix[0]?.length ?? 0;
1215
+ if (columns === 0)
1216
+ throw new Error("Input data must contain at least one feature");
1217
+ for (let i = 0;i < matrix.length; i++) {
1218
+ const row = matrix[i];
1219
+ if (row.length !== columns) {
1220
+ throw new Error(`Row ${i} has ${row.length} columns, expected ${columns}`);
1221
+ }
1222
+ if (!row.every(Number.isFinite)) {
1223
+ throw new Error("Model inputs must contain only finite numbers");
1224
+ }
1225
+ }
1226
+ return matrix;
1227
+ }
1228
+ function validateTargets(y, expectedRows) {
1229
+ if (y.length !== expectedRows) {
1230
+ throw new Error(`X has ${expectedRows} rows but y has ${y.length} elements`);
1231
+ }
1232
+ if (!y.every(Number.isFinite)) {
1233
+ throw new Error("Model targets must contain only finite numbers");
1234
+ }
1235
+ }
1236
+ function validateFeatureCount(X, expectedFeatures) {
1237
+ const actualFeatures = X[0].length;
1238
+ if (actualFeatures !== expectedFeatures) {
1239
+ throw new Error(`Prediction data has ${actualFeatures} features, expected ${expectedFeatures}`);
1240
+ }
1241
+ }
1242
+ function validateFeatureVector(point, expectedFeatures) {
1243
+ if (point.length !== expectedFeatures) {
1244
+ throw new Error(`Prediction point has ${point.length} features, expected ${expectedFeatures}`);
1245
+ }
1246
+ if (!point.every(Number.isFinite)) {
1247
+ throw new Error("Prediction point must contain only finite numbers");
1248
+ }
1249
+ }
1250
+
1199
1251
  // src/models/base.ts
1200
1252
  class BaseRegression {
1201
1253
  _coefficients = [];
1202
1254
  _intercept = 0;
1203
1255
  _fitted = false;
1204
1256
  _fitIntercept;
1257
+ _nFeatures = 0;
1205
1258
  _X = [];
1206
1259
  _y = [];
1207
1260
  _yHat = [];
@@ -1316,22 +1369,34 @@ class BaseRegression {
1316
1369
  `);
1317
1370
  }
1318
1371
  normalizeInput(X) {
1319
- if (X.length === 0)
1320
- throw new Error("Input data cannot be empty");
1321
- if (typeof X[0] === "number") {
1322
- return X.map((v) => [v]);
1323
- }
1324
- return X;
1372
+ return normalizeModelInput(X);
1325
1373
  }
1326
1374
  addInterceptColumn(X) {
1327
1375
  return X.map((row) => [1, ...row]);
1328
1376
  }
1329
1377
  validateFitInput(X, y) {
1330
- if (X.length !== y.length) {
1331
- throw new Error(`X has ${X.length} rows but y has ${y.length} elements`);
1332
- }
1333
- if (X.length === 0)
1334
- throw new Error("Input data cannot be empty");
1378
+ validateTargets(y, X.length);
1379
+ this._nFeatures = X[0].length;
1380
+ }
1381
+ validatePredictInput(X) {
1382
+ this.assertFitted();
1383
+ const Xmat = this.normalizeInput(X);
1384
+ validateFeatureCount(Xmat, this._nFeatures);
1385
+ return Xmat;
1386
+ }
1387
+ predictLinearRows(X) {
1388
+ return X.map((row) => {
1389
+ let sum = this._intercept;
1390
+ for (let j = 0;j < this._coefficients.length; j++) {
1391
+ sum += row[j] * this._coefficients[j];
1392
+ }
1393
+ return sum;
1394
+ });
1395
+ }
1396
+ completeFit(X, statisticsX = X) {
1397
+ this._fitted = true;
1398
+ this._yHat = this.predict(X);
1399
+ this._X = statisticsX;
1335
1400
  }
1336
1401
  assertFitted() {
1337
1402
  if (!this._fitted) {
@@ -1377,19 +1442,11 @@ class LinearRegression extends BaseRegression {
1377
1442
  this._intercept = 0;
1378
1443
  this._coefficients = betaArray;
1379
1444
  }
1380
- this._yHat = this.predict(Xmat);
1381
- this._fitted = true;
1445
+ this.completeFit(Xmat);
1382
1446
  return this;
1383
1447
  }
1384
1448
  predict(X) {
1385
- const Xmat = this.normalizeInput(X);
1386
- return Xmat.map((row) => {
1387
- let sum = this._intercept;
1388
- for (let j = 0;j < this._coefficients.length; j++) {
1389
- sum += row[j] * this._coefficients[j];
1390
- }
1391
- return sum;
1392
- });
1449
+ return this.predictLinearRows(this.validatePredictInput(X));
1393
1450
  }
1394
1451
  }
1395
1452
 
@@ -1427,15 +1484,15 @@ function correlationMatrix(X) {
1427
1484
  xFlat[i * p + j] = X[i][j];
1428
1485
  const wasmResult = engineCorrelationMatrix(xFlat, n, p);
1429
1486
  if (wasmResult) {
1430
- const corr2 = [];
1487
+ const corr = [];
1431
1488
  for (let j1 = 0;j1 < p; j1++) {
1432
1489
  const row = [];
1433
1490
  for (let j2 = 0;j2 < p; j2++) {
1434
1491
  row.push(wasmResult[j1 * p + j2]);
1435
1492
  }
1436
- corr2.push(row);
1493
+ corr.push(row);
1437
1494
  }
1438
- return corr2;
1495
+ return corr;
1439
1496
  }
1440
1497
  const means = [];
1441
1498
  for (let j = 0;j < p; j++) {
@@ -1566,7 +1623,7 @@ function shapiroWilk(data) {
1566
1623
  const n = data.length;
1567
1624
  if (n < 3)
1568
1625
  throw new Error("Shapiro-Wilk requires at least 3 observations");
1569
- const sorted = [...data].sort((a2, b) => a2 - b);
1626
+ const sorted = [...data].sort((a, b) => a - b);
1570
1627
  let mean = 0;
1571
1628
  for (const x of sorted)
1572
1629
  mean += x;
@@ -1625,8 +1682,7 @@ class LassoRegression extends BaseRegression {
1625
1682
  if (wasmResult) {
1626
1683
  this._intercept = wasmResult[0];
1627
1684
  this._coefficients = Array.from(wasmResult.slice(1));
1628
- this._yHat = this.predict(Xmat);
1629
- this._fitted = true;
1685
+ this.completeFit(Xmat);
1630
1686
  return this;
1631
1687
  }
1632
1688
  const { Xstd, xMeans, xStds, yMean } = this.standardize(Xmat, y);
@@ -1666,22 +1722,14 @@ class LassoRegression extends BaseRegression {
1666
1722
  return std > 0.000000000000001 ? bj / std : 0;
1667
1723
  });
1668
1724
  this._intercept = this._fitIntercept ? yMean - this._coefficients.reduce((sum, bj, j) => sum + bj * xMeans[j], 0) : 0;
1669
- this._yHat = this.predict(Xmat);
1670
- this._fitted = true;
1725
+ this.completeFit(Xmat);
1671
1726
  return this;
1672
1727
  }
1673
1728
  getL1Ratio() {
1674
1729
  return 1;
1675
1730
  }
1676
1731
  predict(X) {
1677
- const Xmat = this.normalizeInput(X);
1678
- return Xmat.map((row) => {
1679
- let sum = this._intercept;
1680
- for (let j = 0;j < this._coefficients.length; j++) {
1681
- sum += row[j] * this._coefficients[j];
1682
- }
1683
- return sum;
1684
- });
1732
+ return this.predictLinearRows(this.validatePredictInput(X));
1685
1733
  }
1686
1734
  coordinateUpdate(rho, colNormSq, n, _j) {
1687
1735
  return this.softThreshold(rho, n * this._alpha) / colNormSq;
@@ -1753,28 +1801,29 @@ class KNearestNeighbors {
1753
1801
  _fitted = false;
1754
1802
  _X = [];
1755
1803
  _y = [];
1804
+ _nFeatures = 0;
1756
1805
  constructor(options = {}) {
1757
1806
  this._k = options.k ?? 5;
1758
1807
  this._distance = options.distance ?? "euclidean";
1759
1808
  this._mode = options.mode ?? "classification";
1760
1809
  }
1761
1810
  fit(X, y) {
1762
- const Xmat = this.normalizeInput(X);
1763
- if (Xmat.length !== y.length) {
1764
- throw new Error(`X has ${Xmat.length} rows but y has ${y.length} elements`);
1765
- }
1811
+ const Xmat = normalizeModelInput(X);
1812
+ validateTargets(y, Xmat.length);
1766
1813
  if (Xmat.length < this._k) {
1767
1814
  throw new Error(`Need at least k=${this._k} samples, got ${Xmat.length}`);
1768
1815
  }
1769
1816
  this._X = Xmat;
1770
1817
  this._y = y;
1818
+ this._nFeatures = Xmat[0].length;
1771
1819
  this._fitted = true;
1772
1820
  return this;
1773
1821
  }
1774
1822
  predict(X) {
1775
1823
  if (!this._fitted)
1776
1824
  throw new Error("Model has not been fitted. Call fit() first.");
1777
- const Xmat = this.normalizeInput(X);
1825
+ const Xmat = normalizeModelInput(X);
1826
+ validateFeatureCount(Xmat, this._nFeatures);
1778
1827
  const dim = this._X[0].length;
1779
1828
  const nTrain = this._X.length;
1780
1829
  const nTest = Xmat.length;
@@ -1789,7 +1838,7 @@ class KNearestNeighbors {
1789
1838
  const distMatrix = this._distance === "manhattan" ? engineManhattanDistances(trainFlat, testFlat, nTrain, nTest, dim) : engineEuclideanDistances(trainFlat, testFlat, nTrain, nTest, dim);
1790
1839
  if (distMatrix) {
1791
1840
  return Array.from({ length: nTest }, (_, i) => {
1792
- const indices = Array.from({ length: nTrain }, (_2, j) => j).sort((a, b) => distMatrix[i * nTrain + a] - distMatrix[i * nTrain + b]).slice(0, this._k);
1841
+ const indices = Array.from({ length: nTrain }, (_, j) => j).sort((a, b) => distMatrix[i * nTrain + a] - distMatrix[i * nTrain + b]).slice(0, this._k);
1793
1842
  return this.voteOrMean(indices);
1794
1843
  });
1795
1844
  }
@@ -1798,6 +1847,7 @@ class KNearestNeighbors {
1798
1847
  neighbors(point) {
1799
1848
  if (!this._fitted)
1800
1849
  throw new Error("Model has not been fitted. Call fit() first.");
1850
+ validateFeatureVector(point, this._nFeatures);
1801
1851
  return this.findNeighbors(point).map((n) => n.index);
1802
1852
  }
1803
1853
  voteOrMean(indices) {
@@ -1847,14 +1897,6 @@ class KNearestNeighbors {
1847
1897
  }
1848
1898
  return this._distance === "manhattan" ? sum : Math.sqrt(sum);
1849
1899
  }
1850
- normalizeInput(X) {
1851
- if (X.length === 0)
1852
- throw new Error("Input data cannot be empty");
1853
- if (typeof X[0] === "number") {
1854
- return X.map((v) => [v]);
1855
- }
1856
- return X;
1857
- }
1858
1900
  }
1859
1901
  // src/models/logistic-regression.ts
1860
1902
  class LogisticRegression {
@@ -1866,6 +1908,7 @@ class LogisticRegression {
1866
1908
  _tolerance;
1867
1909
  _y = [];
1868
1910
  _probabilities = [];
1911
+ _nFeatures = 0;
1869
1912
  constructor(options = {}) {
1870
1913
  this._fitIntercept = options.fitIntercept ?? true;
1871
1914
  this._maxIterations = options.maxIterations ?? 100;
@@ -1888,10 +1931,9 @@ class LogisticRegression {
1888
1931
  return expZ / (1 + expZ);
1889
1932
  }
1890
1933
  fit(X, y) {
1891
- const Xmat = this.normalizeInput(X);
1892
- if (Xmat.length !== y.length) {
1893
- throw new Error(`X has ${Xmat.length} rows but y has ${y.length} elements`);
1894
- }
1934
+ const Xmat = normalizeModelInput(X);
1935
+ validateTargets(y, Xmat.length);
1936
+ this._nFeatures = Xmat[0].length;
1895
1937
  for (const yi of y) {
1896
1938
  if (yi !== 0 && yi !== 1) {
1897
1939
  throw new Error("Logistic regression requires binary y (0 or 1)");
@@ -1907,16 +1949,16 @@ class LogisticRegression {
1907
1949
  xFlat[i * k + j] = Xdesign[i][j];
1908
1950
  const wasmBeta = engineIRLSLogistic(xFlat, new Float64Array(y), n, k, this._maxIterations, this._tolerance);
1909
1951
  if (wasmBeta) {
1910
- const betaArray2 = Array.from(wasmBeta);
1952
+ const betaArray = Array.from(wasmBeta);
1911
1953
  if (this._fitIntercept) {
1912
- this._intercept = betaArray2[0];
1913
- this._coefficients = betaArray2.slice(1);
1954
+ this._intercept = betaArray[0];
1955
+ this._coefficients = betaArray.slice(1);
1914
1956
  } else {
1915
1957
  this._intercept = 0;
1916
- this._coefficients = betaArray2;
1958
+ this._coefficients = betaArray;
1917
1959
  }
1918
- this._probabilities = this.predictProbability(Xmat);
1919
1960
  this._fitted = true;
1961
+ this._probabilities = this.predictProbability(Xmat);
1920
1962
  return this;
1921
1963
  }
1922
1964
  const beta = new Float64Array(k);
@@ -1965,15 +2007,18 @@ class LogisticRegression {
1965
2007
  this._intercept = 0;
1966
2008
  this._coefficients = betaArray;
1967
2009
  }
1968
- this._probabilities = this.predictProbability(Xmat);
1969
2010
  this._fitted = true;
2011
+ this._probabilities = this.predictProbability(Xmat);
1970
2012
  return this;
1971
2013
  }
1972
2014
  predict(X) {
1973
2015
  return this.predictProbability(X).map((p) => p >= 0.5 ? 1 : 0);
1974
2016
  }
1975
2017
  predictProbability(X) {
1976
- const Xmat = this.normalizeInput(X);
2018
+ if (!this._fitted)
2019
+ throw new Error("Model has not been fitted. Call fit() first.");
2020
+ const Xmat = normalizeModelInput(X);
2021
+ validateFeatureCount(Xmat, this._nFeatures);
1977
2022
  return Xmat.map((row) => {
1978
2023
  let sum = this._intercept;
1979
2024
  for (let j = 0;j < this._coefficients.length; j++) {
@@ -2037,14 +2082,6 @@ class LogisticRegression {
2037
2082
  bic
2038
2083
  };
2039
2084
  }
2040
- normalizeInput(X) {
2041
- if (X.length === 0)
2042
- throw new Error("Input data cannot be empty");
2043
- if (typeof X[0] === "number") {
2044
- return X.map((v) => [v]);
2045
- }
2046
- return X;
2047
- }
2048
2085
  addInterceptColumn(X) {
2049
2086
  return X.map((row) => [1, ...row]);
2050
2087
  }
@@ -2061,6 +2098,7 @@ class MulticlassLogisticRegression {
2061
2098
  _classes = [];
2062
2099
  _X = [];
2063
2100
  _y = [];
2101
+ _nFeatures = 0;
2064
2102
  constructor(options = {}) {
2065
2103
  this._fitIntercept = options.fitIntercept ?? true;
2066
2104
  this._maxIterations = options.maxIterations ?? 200;
@@ -2084,10 +2122,9 @@ class MulticlassLogisticRegression {
2084
2122
  return exps.map((e) => e / sum);
2085
2123
  }
2086
2124
  fit(X, y) {
2087
- const Xmat = this.normalizeInput(X);
2088
- if (Xmat.length !== y.length) {
2089
- throw new Error(`X has ${Xmat.length} rows but y has ${y.length} elements`);
2090
- }
2125
+ const Xmat = normalizeModelInput(X);
2126
+ validateTargets(y, Xmat.length);
2127
+ this._nFeatures = Xmat[0].length;
2091
2128
  this._X = Xmat;
2092
2129
  this._y = y;
2093
2130
  this._classes = [...new Set(y)].sort((a, b) => a - b);
@@ -2156,7 +2193,8 @@ class MulticlassLogisticRegression {
2156
2193
  predictProbability(X) {
2157
2194
  if (!this._fitted)
2158
2195
  throw new Error("Model has not been fitted. Call fit() first.");
2159
- const Xmat = this.normalizeInput(X);
2196
+ const Xmat = normalizeModelInput(X);
2197
+ validateFeatureCount(Xmat, this._nFeatures);
2160
2198
  const Xdesign = this._fitIntercept ? Xmat.map((row) => [1, ...row]) : Xmat;
2161
2199
  const XMat = Matrix.fromArray(Xdesign);
2162
2200
  const scores = XMat.multiply(this._weights);
@@ -2213,14 +2251,6 @@ class MulticlassLogisticRegression {
2213
2251
  logLikelihood: logLik
2214
2252
  };
2215
2253
  }
2216
- normalizeInput(X) {
2217
- if (X.length === 0)
2218
- throw new Error("Input data cannot be empty");
2219
- if (typeof X[0] === "number") {
2220
- return X.map((v) => [v]);
2221
- }
2222
- return X;
2223
- }
2224
2254
  }
2225
2255
  // src/models/neural-network.ts
2226
2256
  class NeuralNetwork {
@@ -2231,6 +2261,7 @@ class NeuralNetwork {
2231
2261
  _fitted = false;
2232
2262
  _outputSize = 0;
2233
2263
  _classes = [];
2264
+ _nFeatures = 0;
2234
2265
  constructor(options) {
2235
2266
  this._learningRate = options.learningRate ?? 0.01;
2236
2267
  this._epochs = options.epochs ?? 100;
@@ -2263,11 +2294,10 @@ class NeuralNetwork {
2263
2294
  };
2264
2295
  }
2265
2296
  fit(X, y) {
2266
- const Xmat = this.normalizeInput(X);
2267
- if (Xmat.length !== y.length) {
2268
- throw new Error(`X has ${Xmat.length} rows but y has ${y.length} elements`);
2269
- }
2297
+ const Xmat = normalizeModelInput(X);
2298
+ validateTargets(y, Xmat.length);
2270
2299
  const inputSize = Xmat[0].length;
2300
+ this._nFeatures = inputSize;
2271
2301
  const n = Xmat.length;
2272
2302
  if (this._task === "classification") {
2273
2303
  this._classes = [...new Set(y)].sort((a, b) => a - b);
@@ -2295,7 +2325,8 @@ class NeuralNetwork {
2295
2325
  predict(X) {
2296
2326
  if (!this._fitted)
2297
2327
  throw new Error("Model has not been fitted. Call fit() first.");
2298
- const Xmat = this.normalizeInput(X);
2328
+ const Xmat = normalizeModelInput(X);
2329
+ validateFeatureCount(Xmat, this._nFeatures);
2299
2330
  return Xmat.map((row) => {
2300
2331
  const activations = this.forward(row);
2301
2332
  const output = activations[activations.length - 1];
@@ -2316,7 +2347,8 @@ class NeuralNetwork {
2316
2347
  predictRaw(X) {
2317
2348
  if (!this._fitted)
2318
2349
  throw new Error("Model has not been fitted. Call fit() first.");
2319
- const Xmat = this.normalizeInput(X);
2350
+ const Xmat = normalizeModelInput(X);
2351
+ validateFeatureCount(Xmat, this._nFeatures);
2320
2352
  return Xmat.map((row) => {
2321
2353
  const activations = this.forward(row);
2322
2354
  return Array.from(activations[activations.length - 1]);
@@ -2432,14 +2464,6 @@ class NeuralNetwork {
2432
2464
  encoded[idx] = 1;
2433
2465
  return encoded;
2434
2466
  }
2435
- normalizeInput(X) {
2436
- if (X.length === 0)
2437
- throw new Error("Input data cannot be empty");
2438
- if (typeof X[0] === "number") {
2439
- return X.map((v) => [v]);
2440
- }
2441
- return X;
2442
- }
2443
2467
  }
2444
2468
  // src/models/polynomial-regression.ts
2445
2469
  class PolynomialRegression extends BaseRegression {
@@ -2472,12 +2496,11 @@ class PolynomialRegression extends BaseRegression {
2472
2496
  this._inner.fit(Xexpanded, y);
2473
2497
  this._coefficients = this._inner.coefficients;
2474
2498
  this._intercept = this._inner.intercept;
2475
- this._yHat = this.predict(Xmat);
2476
- this._fitted = true;
2499
+ this.completeFit(Xmat, Xexpanded);
2477
2500
  return this;
2478
2501
  }
2479
2502
  predict(X) {
2480
- const Xmat = this.normalizeInput(X);
2503
+ const Xmat = this.validatePredictInput(X);
2481
2504
  const Xexpanded = this.expandFeatures(Xmat);
2482
2505
  return this._inner.predict(Xexpanded);
2483
2506
  }
@@ -2514,19 +2537,11 @@ class RidgeRegression extends BaseRegression {
2514
2537
  this._intercept = 0;
2515
2538
  this._coefficients = betaArray;
2516
2539
  }
2517
- this._yHat = this.predict(Xmat);
2518
- this._fitted = true;
2540
+ this.completeFit(Xmat);
2519
2541
  return this;
2520
2542
  }
2521
2543
  predict(X) {
2522
- const Xmat = this.normalizeInput(X);
2523
- return Xmat.map((row) => {
2524
- let sum = this._intercept;
2525
- for (let j = 0;j < this._coefficients.length; j++) {
2526
- sum += row[j] * this._coefficients[j];
2527
- }
2528
- return sum;
2529
- });
2544
+ return this.predictLinearRows(this.validatePredictInput(X));
2530
2545
  }
2531
2546
  }
2532
2547
  // src/models/weighted-regression.ts
@@ -2563,19 +2578,11 @@ class WeightedRegression extends BaseRegression {
2563
2578
  this._intercept = 0;
2564
2579
  this._coefficients = betaArray;
2565
2580
  }
2566
- this._yHat = this.predict(Xmat);
2567
- this._fitted = true;
2581
+ this.completeFit(Xmat);
2568
2582
  return this;
2569
2583
  }
2570
2584
  predict(X) {
2571
- const Xmat = this.normalizeInput(X);
2572
- return Xmat.map((row) => {
2573
- let sum = this._intercept;
2574
- for (let j = 0;j < this._coefficients.length; j++) {
2575
- sum += row[j] * this._coefficients[j];
2576
- }
2577
- return sum;
2578
- });
2585
+ return this.predictLinearRows(this.validatePredictInput(X));
2579
2586
  }
2580
2587
  }
2581
2588
 
@@ -2604,7 +2611,7 @@ class RobustRegression extends BaseRegression {
2604
2611
  for (let iter = 0;iter < this._maxIterations; iter++) {
2605
2612
  const prevCoeffs = [...this._coefficients];
2606
2613
  const prevIntercept = this._intercept;
2607
- const yHat = this.predict(Xmat);
2614
+ const yHat = this.predictLinearRows(Xmat);
2608
2615
  const residuals = y.map((yi, i) => yi - yHat[i]);
2609
2616
  const sortedResiduals = [...residuals].sort((a, b) => a - b);
2610
2617
  const residualMedian = median(sortedResiduals);
@@ -2626,19 +2633,11 @@ class RobustRegression extends BaseRegression {
2626
2633
  if (maxChange < this._tolerance)
2627
2634
  break;
2628
2635
  }
2629
- this._yHat = this.predict(Xmat);
2630
- this._fitted = true;
2636
+ this.completeFit(Xmat);
2631
2637
  return this;
2632
2638
  }
2633
2639
  predict(X) {
2634
- const Xmat = this.normalizeInput(X);
2635
- return Xmat.map((row) => {
2636
- let sum = this._intercept;
2637
- for (let j = 0;j < this._coefficients.length; j++) {
2638
- sum += row[j] * this._coefficients[j];
2639
- }
2640
- return sum;
2641
- });
2640
+ return this.predictLinearRows(this.validatePredictInput(X));
2642
2641
  }
2643
2642
  }
2644
2643
  function huberWeight(u, k) {