Skip to content

Commit 4d632c9

Browse files
committed
add option for selecting gradient descent method for fitting linear regression
1 parent 167fc95 commit 4d632c9

2 files changed

Lines changed: 30 additions & 4 deletions

File tree

lib/learn_kit/regression/linear.ex

Lines changed: 25 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -50,11 +50,16 @@ defmodule LearnKit.Regression.Linear do
5050
end
5151

5252
@doc """
53-
Fit train data
53+
Fit train data with least squares method
5454
5555
## Parameters
5656
5757
- predictor: %LearnKit.Regression.Linear{}
58+
- options: keyword list with options
59+
60+
## Options
61+
62+
- method: method for fit, "least squares"/"gradient descent", default is "least squares", optional
5863
5964
## Examples
6065
@@ -65,11 +70,28 @@ defmodule LearnKit.Regression.Linear do
6570
results: [3, 6, 10, 15]
6671
}
6772
73+
iex> predictor = predictor |> LearnKit.Regression.Linear.fit([method: "gradient descent"])
74+
%LearnKit.Regression.Linear{
75+
coefficients: [-1.5, 4.0],
76+
factors: [1, 2, 3, 4],
77+
results: [3, 6, 10, 15]
78+
}
79+
6880
"""
6981
@spec fit(%Linear{factors: factors, results: results}) :: %Linear{factors: factors, results: results, coefficients: coefficients}
7082

71-
def fit(%Linear{factors: factors, results: results}) do
72-
%Linear{factors: factors, results: results, coefficients: fit_data(factors, results)}
83+
def fit(%Linear{factors: factors, results: results}, options \\ []) do
84+
coefficients = Keyword.merge([method: ""], options)
85+
|> define_method_for_fit
86+
|> fit_data(factors, results)
87+
%Linear{factors: factors, results: results, coefficients: coefficients}
88+
end
89+
90+
defp define_method_for_fit(options) do
91+
cond do
92+
Keyword.get(options, :method) == "gradient descent" -> "gradient descent"
93+
true -> ""
94+
end
7395
end
7496

7597
@doc """

lib/learn_kit/regression/linear/fit.ex

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,11 @@ defmodule LearnKit.Regression.Linear.Fit do
77

88
defmacro __using__(_opts) do
99
quote do
10-
defp fit_data(factors, results) do
10+
defp fit_data(method, factors, results) when method == "gradient descent" do
11+
12+
end
13+
14+
defp fit_data(_, factors, results) do
1115
beta = Math.correlation(factors, results) * Math.standard_deviation(results) / Math.standard_deviation(factors)
1216
alpha = Math.mean(results) - beta * Math.mean(factors)
1317
[alpha, beta]

0 commit comments

Comments
 (0)