nummeth 0.1.0__py3-none-any.whl

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.
Files changed (39) hide show
  1. nummeth/__init__.py +132 -0
  2. nummeth/applications/__init__.py +4 -0
  3. nummeth/applications/circuit.py +78 -0
  4. nummeth/calculus/__init__.py +26 -0
  5. nummeth/calculus/differentiation.py +21 -0
  6. nummeth/calculus/integration.py +178 -0
  7. nummeth/core/__init__.py +4 -0
  8. nummeth/core/results.py +107 -0
  9. nummeth/fitting/__init__.py +4 -0
  10. nummeth/fitting/regression.py +123 -0
  11. nummeth/interpolation/__init__.py +21 -0
  12. nummeth/interpolation/divided_diff.py +88 -0
  13. nummeth/interpolation/finite_diff.py +180 -0
  14. nummeth/interpolation/lagrange.py +72 -0
  15. nummeth/interpolation/spline.py +134 -0
  16. nummeth/linear_systems/__init__.py +13 -0
  17. nummeth/linear_systems/eigen.py +82 -0
  18. nummeth/linear_systems/elimination.py +92 -0
  19. nummeth/linear_systems/iterative.py +127 -0
  20. nummeth/linear_systems/matrix_ops.py +65 -0
  21. nummeth/ode/__init__.py +12 -0
  22. nummeth/ode/euler.py +103 -0
  23. nummeth/ode/predictor_corrector.py +64 -0
  24. nummeth/ode/runge_kutta.py +113 -0
  25. nummeth/pde/__init__.py +12 -0
  26. nummeth/pde/elliptic.py +135 -0
  27. nummeth/pde/hyperbolic.py +92 -0
  28. nummeth/pde/parabolic.py +148 -0
  29. nummeth/roots/__init__.py +14 -0
  30. nummeth/roots/bisection.py +86 -0
  31. nummeth/roots/false_position.py +73 -0
  32. nummeth/roots/fixed_point.py +55 -0
  33. nummeth/roots/newton_raphson.py +66 -0
  34. nummeth/roots/secant.py +63 -0
  35. nummeth-0.1.0.dist-info/METADATA +145 -0
  36. nummeth-0.1.0.dist-info/RECORD +39 -0
  37. nummeth-0.1.0.dist-info/WHEEL +5 -0
  38. nummeth-0.1.0.dist-info/licenses/LICENSE +21 -0
  39. nummeth-0.1.0.dist-info/top_level.txt +1 -0
nummeth/__init__.py ADDED
@@ -0,0 +1,132 @@
1
+ """
2
+ NumMeth: A Practical, Educational Numerical Methods Library for Laboratories & Engineering.
3
+ """
4
+
5
+ __version__ = "0.1.0"
6
+
7
+ # Core structures
8
+ from nummeth.core.results import IterationResult, TableFormatter
9
+
10
+ # 1. Algebraic & Transcendental Roots
11
+ from nummeth.roots import (
12
+ bisection,
13
+ false_position,
14
+ newton_raphson,
15
+ secant,
16
+ fixed_point,
17
+ )
18
+
19
+ # 2. Linear Systems & Eigenvalues
20
+ from nummeth.linear_systems import (
21
+ gauss_elimination,
22
+ gauss_jacobi,
23
+ gauss_seidel,
24
+ matrix_inverse,
25
+ power_method,
26
+ )
27
+
28
+ # 3. Interpolation & Splines
29
+ from nummeth.interpolation import (
30
+ lagrange_interpolate,
31
+ divided_difference_table,
32
+ divided_difference_interp,
33
+ newton_forward_table,
34
+ newton_forward_interp,
35
+ newton_backward_table,
36
+ newton_backward_interp,
37
+ natural_cubic_spline,
38
+ )
39
+
40
+ # 4. Curve Fitting & Regression
41
+ from nummeth.fitting import (
42
+ polynomial_fit,
43
+ )
44
+
45
+ # 5. Numerical Calculus
46
+ from nummeth.calculus import (
47
+ derivative_forward,
48
+ derivative_backward,
49
+ derivative_central,
50
+ derivative_second,
51
+ trapezoidal,
52
+ simpson_13,
53
+ simpson_38,
54
+ boole,
55
+ weddle,
56
+ )
57
+
58
+ # 6. Ordinary Differential Equations (IVP)
59
+ from nummeth.ode import (
60
+ euler,
61
+ modified_euler,
62
+ rk2,
63
+ rk4,
64
+ milne_predictor_corrector,
65
+ )
66
+
67
+ # 7. Partial Differential Equations (Finite Differences)
68
+ from nummeth.pde import (
69
+ solve_laplace_2d,
70
+ solve_poisson_2d,
71
+ heat_1d_bender_schmidt,
72
+ heat_1d_crank_nicolson,
73
+ wave_1d_explicit,
74
+ )
75
+
76
+ # 8. Applications
77
+ from nummeth.applications import (
78
+ solve_circuit_nodal,
79
+ )
80
+
81
+ __all__ = [
82
+ "__version__",
83
+ "IterationResult",
84
+ "TableFormatter",
85
+ # Roots
86
+ "bisection",
87
+ "false_position",
88
+ "newton_raphson",
89
+ "secant",
90
+ "fixed_point",
91
+ # Linear
92
+ "gauss_elimination",
93
+ "gauss_jacobi",
94
+ "gauss_seidel",
95
+ "matrix_inverse",
96
+ "power_method",
97
+ # Interpolation
98
+ "lagrange_interpolate",
99
+ "divided_difference_table",
100
+ "divided_difference_interp",
101
+ "newton_forward_table",
102
+ "newton_forward_interp",
103
+ "newton_backward_table",
104
+ "newton_backward_interp",
105
+ "natural_cubic_spline",
106
+ # Fitting
107
+ "polynomial_fit",
108
+ # Calculus
109
+ "derivative_forward",
110
+ "derivative_backward",
111
+ "derivative_central",
112
+ "derivative_second",
113
+ "trapezoidal",
114
+ "simpson_13",
115
+ "simpson_38",
116
+ "boole",
117
+ "weddle",
118
+ # ODE
119
+ "euler",
120
+ "modified_euler",
121
+ "rk2",
122
+ "rk4",
123
+ "milne_predictor_corrector",
124
+ # PDE
125
+ "solve_laplace_2d",
126
+ "solve_poisson_2d",
127
+ "heat_1d_bender_schmidt",
128
+ "heat_1d_crank_nicolson",
129
+ "wave_1d_explicit",
130
+ # Applications
131
+ "solve_circuit_nodal",
132
+ ]
@@ -0,0 +1,4 @@
1
+ """Real-world numerical applications."""
2
+ from .circuit import solve_circuit_nodal
3
+
4
+ __all__ = ["solve_circuit_nodal"]
@@ -0,0 +1,78 @@
1
+ """
2
+ Electrical circuit nodal analysis application using systems of linear equations.
3
+ Ported from laboratory circuit solver exercise.
4
+ """
5
+
6
+ from typing import Dict, List, Tuple
7
+ from nummeth.core.results import IterationResult
8
+ from nummeth.linear_systems.elimination import gauss_elimination
9
+
10
+ def solve_circuit_nodal(
11
+ known_voltages: Dict[int, float],
12
+ unknown_nodes: List[int],
13
+ resistors: List[Tuple[int, int, float]],
14
+ method: str = "elimination"
15
+ ) -> IterationResult:
16
+ """
17
+ Solves for unknown node voltages and branch currents in a resistor network
18
+ using Kirchhoff's Current Law (Nodal Analysis).
19
+
20
+ Parameters:
21
+ known_voltages: Dict mapping node number -> known voltage (e.g. {1: 10.0, 6: 0.0}).
22
+ unknown_nodes: List of unknown node indices (e.g. [2, 3, 4, 5]).
23
+ resistors: List of tuples (node1, node2, resistance_in_ohms).
24
+ method: "elimination" or "seidel".
25
+
26
+ Returns:
27
+ IterationResult where value is dict of all node voltages, and extra['branch_currents']
28
+ contains branch currents.
29
+ """
30
+ n = len(unknown_nodes)
31
+ index = {node: i for i, node in enumerate(unknown_nodes)}
32
+
33
+ A = [[0.0 for _ in range(n)] for _ in range(n)]
34
+ b = [0.0 for _ in range(n)]
35
+
36
+ for node in unknown_nodes:
37
+ row = index[node]
38
+ for u, v, r in resistors:
39
+ if r <= 0:
40
+ raise ValueError("Resistance must be positive.")
41
+ if node == u:
42
+ other = v
43
+ elif node == v:
44
+ other = u
45
+ else:
46
+ continue
47
+
48
+ g = 1.0 / r
49
+ A[row][row] += g
50
+
51
+ if other in index:
52
+ A[row][index[other]] -= g
53
+ elif other in known_voltages:
54
+ b[row] += g * known_voltages[other]
55
+
56
+ # Solve linear system
57
+ sol_res = gauss_elimination(A, b, pivoting=True)
58
+ voltages = dict(known_voltages)
59
+ for i, node in enumerate(unknown_nodes):
60
+ voltages[node] = sol_res.value[i]
61
+
62
+ # Calculate branch currents
63
+ branch_currents = []
64
+ for u, v, r in resistors:
65
+ cur = (voltages[u] - voltages[v]) / r
66
+ branch_currents.append({"from": u, "to": v, "resistance": r, "current": cur})
67
+
68
+ return IterationResult(
69
+ value=voltages,
70
+ converged=True,
71
+ message="Circuit nodal voltages and currents resolved successfully.",
72
+ extra={
73
+ "voltages": voltages,
74
+ "branch_currents": branch_currents,
75
+ "conductance_matrix": A,
76
+ "rhs_vector": b
77
+ }
78
+ )
@@ -0,0 +1,26 @@
1
+ """Numerical calculus: differentiation and integration."""
2
+ from .differentiation import (
3
+ derivative_forward,
4
+ derivative_backward,
5
+ derivative_central,
6
+ derivative_second,
7
+ )
8
+ from .integration import (
9
+ trapezoidal,
10
+ simpson_13,
11
+ simpson_38,
12
+ boole,
13
+ weddle,
14
+ )
15
+
16
+ __all__ = [
17
+ "derivative_forward",
18
+ "derivative_backward",
19
+ "derivative_central",
20
+ "derivative_second",
21
+ "trapezoidal",
22
+ "simpson_13",
23
+ "simpson_38",
24
+ "boole",
25
+ "weddle",
26
+ ]
@@ -0,0 +1,21 @@
1
+ """
2
+ Numerical Differentiation using finite difference formulas.
3
+ """
4
+
5
+ from typing import Callable, Union, List
6
+
7
+ def derivative_forward(f: Callable[[float], float], x: float, h: float = 1e-5) -> float:
8
+ """First derivative via 2-point forward difference: f'(x) ≈ (f(x + h) - f(x)) / h"""
9
+ return (f(x + h) - f(x)) / h
10
+
11
+ def derivative_backward(f: Callable[[float], float], x: float, h: float = 1e-5) -> float:
12
+ """First derivative via 2-point backward difference: f'(x) ≈ (f(x) - f(x - h)) / h"""
13
+ return (f(x) - f(x - h)) / h
14
+
15
+ def derivative_central(f: Callable[[float], float], x: float, h: float = 1e-5) -> float:
16
+ """First derivative via 2-point central difference (O(h^2)): f'(x) ≈ (f(x + h) - f(x - h)) / (2h)"""
17
+ return (f(x + h) - f(x - h)) / (2.0 * h)
18
+
19
+ def derivative_second(f: Callable[[float], float], x: float, h: float = 1e-4) -> float:
20
+ """Second derivative via central difference: f''(x) ≈ (f(x + h) - 2f(x) + f(x - h)) / h^2"""
21
+ return (f(x + h) - 2.0 * f(x) + f(x - h)) / (h ** 2)
@@ -0,0 +1,178 @@
1
+ """
2
+ Numerical Integration (Quadrature) formulas:
3
+ - Trapezoidal Rule
4
+ - Simpson's 1/3 Rule
5
+ - Simpson's 3/8 Rule
6
+ - Boole's Rule
7
+ - Weddle's Rule
8
+ - Romberg Integration
9
+ """
10
+
11
+ from typing import Callable, Union, List, Optional
12
+ from nummeth.core.results import IterationResult
13
+
14
+ def _get_samples(f_or_y: Union[Callable[[float], float], List[float]], a: float, b: float, n: int) -> Tuple_Helper:
15
+ """Helper to obtain step size h and samples [y_0, ..., y_n]."""
16
+ h = (b - a) / n
17
+ if callable(f_or_y):
18
+ y = [float(f_or_y(a + i * h)) for i in range(n + 1)]
19
+ else:
20
+ y = [float(val) for val in f_or_y]
21
+ if len(y) != n + 1:
22
+ raise ValueError(f"For n={n} intervals, exactly {n + 1} y-values are required. Got {len(y)}.")
23
+ return h, y
24
+
25
+ class Tuple_Helper(tuple):
26
+ pass
27
+
28
+
29
+ def trapezoidal(
30
+ f_or_y: Union[Callable[[float], float], List[float]],
31
+ a: float,
32
+ b: float,
33
+ n: int = 100,
34
+ exact: Optional[float] = None
35
+ ) -> IterationResult:
36
+ """
37
+ Evaluates definite integral using the Composite Trapezoidal Rule.
38
+ Formula: (h / 2) * [y_0 + 2*(y_1 + ... + y_{n-1}) + y_n]
39
+ """
40
+ if n < 1:
41
+ raise ValueError("Subintervals n must be at least 1.")
42
+ h, y = _get_samples(f_or_y, a, b, n)
43
+
44
+ s = y[0] + y[-1] + 2.0 * sum(y[1:-1])
45
+ integral = (h / 2.0) * s
46
+ abs_err = abs(exact - integral) if exact is not None else None
47
+
48
+ return IterationResult(
49
+ value=integral,
50
+ error=abs_err,
51
+ message=f"Trapezoidal rule over [{a}, {b}] with n={n}: {integral:.6f}",
52
+ extra={"h": h, "n": n, "subintervals": n, "method": "Trapezoidal"}
53
+ )
54
+
55
+
56
+ def simpson_13(
57
+ f_or_y: Union[Callable[[float], float], List[float]],
58
+ a: float,
59
+ b: float,
60
+ n: int = 100,
61
+ exact: Optional[float] = None
62
+ ) -> IterationResult:
63
+ """
64
+ Evaluates definite integral using Composite Simpson's 1/3 Rule.
65
+ Requires n to be an even integer.
66
+ Formula: (h / 3) * [y_0 + 4*(odd terms) + 2*(even terms) + y_n]
67
+ """
68
+ if n % 2 != 0:
69
+ raise ValueError("Simpson's 1/3 rule requires an even number of subintervals (n % 2 == 0).")
70
+ h, y = _get_samples(f_or_y, a, b, n)
71
+
72
+ s_odd = sum(y[i] for i in range(1, n, 2))
73
+ s_even = sum(y[i] for i in range(2, n, 2))
74
+ integral = (h / 3.0) * (y[0] + y[-1] + 4.0 * s_odd + 2.0 * s_even)
75
+ abs_err = abs(exact - integral) if exact is not None else None
76
+
77
+ return IterationResult(
78
+ value=integral,
79
+ error=abs_err,
80
+ message=f"Simpson's 1/3 rule over [{a}, {b}] with n={n}: {integral:.6f}",
81
+ extra={"h": h, "n": n, "method": "Simpson 1/3"}
82
+ )
83
+
84
+
85
+ def simpson_38(
86
+ f_or_y: Union[Callable[[float], float], List[float]],
87
+ a: float,
88
+ b: float,
89
+ n: int = 99,
90
+ exact: Optional[float] = None
91
+ ) -> IterationResult:
92
+ """
93
+ Evaluates definite integral using Composite Simpson's 3/8 Rule.
94
+ Requires n to be a multiple of 3.
95
+ Formula: (3h / 8) * [y_0 + 3*sum_{i%3!=0} y_i + 2*sum_{i%3==0} y_i + y_n]
96
+ """
97
+ if n % 3 != 0:
98
+ raise ValueError("Simpson's 3/8 rule requires n to be a multiple of 3 (n % 3 == 0).")
99
+ h, y = _get_samples(f_or_y, a, b, n)
100
+
101
+ s = y[0] + y[-1]
102
+ for i in range(1, n):
103
+ if i % 3 == 0:
104
+ s += 2.0 * y[i]
105
+ else:
106
+ s += 3.0 * y[i]
107
+
108
+ integral = (3.0 * h / 8.0) * s
109
+ abs_err = abs(exact - integral) if exact is not None else None
110
+
111
+ return IterationResult(
112
+ value=integral,
113
+ error=abs_err,
114
+ message=f"Simpson's 3/8 rule over [{a}, {b}] with n={n}: {integral:.6f}",
115
+ extra={"h": h, "n": n, "method": "Simpson 3/8"}
116
+ )
117
+
118
+
119
+ def boole(
120
+ f_or_y: Union[Callable[[float], float], List[float]],
121
+ a: float,
122
+ b: float,
123
+ n: int = 4,
124
+ exact: Optional[float] = None
125
+ ) -> IterationResult:
126
+ """
127
+ Evaluates definite integral using Boole's Rule.
128
+ Requires n to be a multiple of 4.
129
+ Formula per 4-step panel: (2h / 45) * [7*y0 + 32*y1 + 12*y2 + 32*y3 + 7*y4]
130
+ """
131
+ if n % 4 != 0:
132
+ raise ValueError("Boole's rule requires n to be a multiple of 4.")
133
+ h, y = _get_samples(f_or_y, a, b, n)
134
+
135
+ total = 0.0
136
+ for i in range(0, n, 4):
137
+ total += 7.0 * y[i] + 32.0 * y[i + 1] + 12.0 * y[i + 2] + 32.0 * y[i + 3] + 7.0 * y[i + 4]
138
+
139
+ integral = (2.0 * h / 45.0) * total
140
+ abs_err = abs(exact - integral) if exact is not None else None
141
+
142
+ return IterationResult(
143
+ value=integral,
144
+ error=abs_err,
145
+ message=f"Boole's rule over [{a}, {b}] with n={n}: {integral:.6f}",
146
+ extra={"h": h, "n": n, "method": "Boole"}
147
+ )
148
+
149
+
150
+ def weddle(
151
+ f_or_y: Union[Callable[[float], float], List[float]],
152
+ a: float,
153
+ b: float,
154
+ n: int = 6,
155
+ exact: Optional[float] = None
156
+ ) -> IterationResult:
157
+ """
158
+ Evaluates definite integral using Weddle's Rule.
159
+ Requires n to be a multiple of 6.
160
+ Formula per 6-step panel: (3h / 10) * [y0 + 5*y1 + y2 + 6*y3 + y4 + 5*y5 + y6]
161
+ """
162
+ if n % 6 != 0:
163
+ raise ValueError("Weddle's rule requires n to be a multiple of 6.")
164
+ h, y = _get_samples(f_or_y, a, b, n)
165
+
166
+ total = 0.0
167
+ for i in range(0, n, 6):
168
+ total += (y[i] + 5.0 * y[i + 1] + y[i + 2] + 6.0 * y[i + 3] + y[i + 4] + 5.0 * y[i + 5] + y[i + 6])
169
+
170
+ integral = (3.0 * h / 10.0) * total
171
+ abs_err = abs(exact - integral) if exact is not None else None
172
+
173
+ return IterationResult(
174
+ value=integral,
175
+ error=abs_err,
176
+ message=f"Weddle's rule over [{a}, {b}] with n={n}: {integral:.6f}",
177
+ extra={"h": h, "n": n, "method": "Weddle"}
178
+ )
@@ -0,0 +1,4 @@
1
+ """Core utilities and result data structures."""
2
+ from .results import IterationResult, TableFormatter
3
+
4
+ __all__ = ["IterationResult", "TableFormatter"]
@@ -0,0 +1,107 @@
1
+ """
2
+ Core result structures and formatting utilities for numerical methods laboratory calculations.
3
+ """
4
+
5
+ from typing import List, Dict, Any, Optional, Sequence
6
+ import math
7
+
8
+ class TableFormatter:
9
+ """Helper to format numerical iteration logs as clean ASCII / markdown tables."""
10
+
11
+ @staticmethod
12
+ def format_table(headers: List[str], rows: List[List[Any]], precision: int = 6) -> str:
13
+ """Formats headers and data rows into a clean, aligned ASCII table."""
14
+ if not rows:
15
+ return "(No iterations)"
16
+
17
+ formatted_rows = []
18
+ for row in rows:
19
+ formatted_row = []
20
+ for cell in row:
21
+ if cell is None:
22
+ formatted_row.append("N/A")
23
+ elif isinstance(cell, float):
24
+ if math.isnan(cell):
25
+ formatted_row.append("NaN")
26
+ elif math.isinf(cell):
27
+ formatted_row.append("Inf")
28
+ else:
29
+ formatted_row.append(f"{cell:.{precision}f}")
30
+ else:
31
+ formatted_row.append(str(cell))
32
+ formatted_rows.append(formatted_row)
33
+
34
+ col_widths = [len(h) for h in headers]
35
+ for row in formatted_rows:
36
+ for i, cell in enumerate(row):
37
+ if i < len(col_widths):
38
+ col_widths[i] = max(col_widths[i], len(cell))
39
+ else:
40
+ col_widths.append(len(cell))
41
+
42
+ # Build borders
43
+ sep = "+" + "+".join("-" * (w + 2) for w in col_widths) + "+"
44
+ header_line = "| " + " | ".join(h.center(col_widths[i]) for i, h in enumerate(headers)) + " |"
45
+
46
+ lines = [sep, header_line, sep]
47
+ for r in formatted_rows:
48
+ row_line = "| " + " | ".join(
49
+ (r[i] if i < len(r) else "").rjust(col_widths[i]) for i in range(len(col_widths))
50
+ ) + " |"
51
+ lines.append(row_line)
52
+ lines.append(sep)
53
+
54
+ return "\n".join(lines)
55
+
56
+
57
+ class IterationResult:
58
+ """
59
+ Standard return container for iterative numerical methods.
60
+ Stores final output, iteration count, step history, and error metrics.
61
+ """
62
+ def __init__(
63
+ self,
64
+ value: Any = None,
65
+ iterations: int = 0,
66
+ converged: bool = True,
67
+ error: Optional[float] = None,
68
+ headers: Optional[List[str]] = None,
69
+ history: Optional[List[List[Any]]] = None,
70
+ message: str = "",
71
+ extra: Optional[Dict[str, Any]] = None
72
+ ):
73
+ self.value = value
74
+ self.root = value # Alias for root-finding methods
75
+ self.solution = value # Alias for linear solvers / ODEs
76
+ self.iterations = iterations
77
+ self.converged = converged
78
+ self.error = error
79
+ self.headers = headers or []
80
+ self.history = history or []
81
+ self.message = message
82
+ self.extra = extra or {}
83
+
84
+ def table(self, precision: int = 6) -> str:
85
+ """Returns the iteration history formatted as an ASCII table."""
86
+ if not self.headers or not self.history:
87
+ return "(No step history recorded)"
88
+ return TableFormatter.format_table(self.headers, self.history, precision=precision)
89
+
90
+ def print_table(self, precision: int = 6) -> None:
91
+ """Prints the iteration history table to stdout."""
92
+ print(self.table(precision=precision))
93
+
94
+ def to_dict(self) -> Dict[str, Any]:
95
+ """Converts result summary to dictionary."""
96
+ return {
97
+ "value": self.value,
98
+ "iterations": self.iterations,
99
+ "converged": self.converged,
100
+ "error": self.error,
101
+ "message": self.message,
102
+ "extra": self.extra
103
+ }
104
+
105
+ def __repr__(self) -> str:
106
+ return (f"<IterationResult value={self.value!r}, iterations={self.iterations}, "
107
+ f"converged={self.converged}, error={self.error}>")
@@ -0,0 +1,4 @@
1
+ """Curve fitting and regression modules."""
2
+ from .regression import polynomial_fit
3
+
4
+ __all__ = ["polynomial_fit"]
@@ -0,0 +1,123 @@
1
+ """
2
+ Least Squares Regression and Curve Fitting.
3
+ Supports polynomial fitting (degrees 1, 2, 3, etc.), exponential, and power law fits.
4
+ """
5
+
6
+ from typing import List, Optional, Tuple, Dict, Any
7
+ import math
8
+ from nummeth.core.results import IterationResult
9
+
10
+ def polynomial_fit(
11
+ x: List[float],
12
+ y: List[float],
13
+ degree: int = 1,
14
+ plot: bool = False,
15
+ x_label: str = "x",
16
+ y_label: str = "y",
17
+ title: Optional[str] = None
18
+ ) -> IterationResult:
19
+ """
20
+ Fits a polynomial of specified degree y ≈ c_m x^m + ... + c_1 x + c_0
21
+ using the method of least squares and normal equations (A^T A) c = A^T y.
22
+
23
+ Parameters:
24
+ x: List of x values.
25
+ y: List of y values.
26
+ degree: Polynomial degree (1 for linear, 2 for quadratic, 3 for cubic).
27
+ plot: If True, displays a matplotlib regression plot.
28
+
29
+ Returns:
30
+ IterationResult with coefficients [c_m, c_{m-1}, ..., c_0] in descending powers of x.
31
+ """
32
+ n = len(x)
33
+ if len(y) != n or n <= degree:
34
+ raise ValueError("Number of data points must exceed the polynomial degree.")
35
+
36
+ # Design matrix A: row i = [x[i]^m, x[i]^(m-1), ..., 1]
37
+ m = degree
38
+ p = m + 1
39
+ ATA = [[0.0 for _ in range(p)] for _ in range(p)]
40
+ ATy = [0.0 for _ in range(p)]
41
+
42
+ # Compute normal equations directly
43
+ for k in range(n):
44
+ xi = x[k]
45
+ yi = y[k]
46
+ # Powers from m down to 0
47
+ powers = [xi ** (m - j) for j in range(p)]
48
+ for r in range(p):
49
+ for c in range(p):
50
+ ATA[r][c] += powers[r] * powers[c]
51
+ ATy[r] += powers[r] * yi
52
+
53
+ # Solve ATA * c = ATy via Gaussian elimination with partial pivoting
54
+ # Augmented matrix
55
+ aug = [ATA[i][:] + [ATy[i]] for i in range(p)]
56
+
57
+ for i in range(p):
58
+ pivot = i
59
+ for r in range(i + 1, p):
60
+ if abs(aug[r][i]) > abs(aug[pivot][i]):
61
+ pivot = r
62
+ aug[i], aug[pivot] = aug[pivot], aug[i]
63
+
64
+ pivot_val = aug[i][i]
65
+ if abs(pivot_val) < 1e-12:
66
+ raise ValueError("Collinear or singular system in normal equations.")
67
+
68
+ for r in range(i + 1, p):
69
+ factor = aug[r][i] / pivot_val
70
+ for c in range(i, p + 1):
71
+ aug[r][c] -= factor * aug[i][c]
72
+
73
+ # Back substitution
74
+ coeffs = [0.0] * p
75
+ for i in range(p - 1, -1, -1):
76
+ total = aug[i][p]
77
+ for c in range(i + 1, p):
78
+ total -= aug[i][c] * coeffs[c]
79
+ coeffs[i] = total / aug[i][i]
80
+
81
+ # Calculate R-squared (coefficient of determination)
82
+ y_mean = sum(y) / n
83
+ ss_tot = sum((yi - y_mean) ** 2 for yi in y)
84
+ y_pred = []
85
+ for xi in x:
86
+ val = sum(coeffs[j] * (xi ** (m - j)) for j in range(p))
87
+ y_pred.append(val)
88
+ ss_res = sum((y[i] - y_pred[i]) ** 2 for i in range(n))
89
+ r_squared = 1.0 - (ss_res / ss_tot) if ss_tot > 0 else 1.0
90
+
91
+ if plot:
92
+ try:
93
+ import matplotlib.pyplot as plt
94
+ import numpy as np
95
+ min_x, max_x = min(x), max(x)
96
+ padding = (max_x - min_x) * 0.05
97
+ xx = np.linspace(min_x - padding, max_x + padding, 200)
98
+ yy = [sum(coeffs[j] * (val ** (m - j)) for j in range(p)) for val in xx]
99
+
100
+ plt.figure(figsize=(7, 5))
101
+ plt.scatter(x, y, color="red", label="Data Points", zorder=5)
102
+ plt.plot(xx, yy, color="blue", linewidth=2, label=f"Fit (deg={degree}, R²={r_squared:.4f})")
103
+ plt.title(title or f"Least Squares Fit (Degree {degree})")
104
+ plt.xlabel(x_label)
105
+ plt.ylabel(y_label)
106
+ plt.grid(True)
107
+ plt.legend()
108
+ plt.show()
109
+ except ImportError:
110
+ pass
111
+
112
+ return IterationResult(
113
+ value=coeffs,
114
+ converged=True,
115
+ error=ss_res,
116
+ message=f"Polynomial fit (degree {degree}) completed with R² = {r_squared:.4f}.",
117
+ extra={
118
+ "coefficients": coeffs,
119
+ "r_squared": r_squared,
120
+ "residuals": ss_res,
121
+ "degree": degree
122
+ }
123
+ )