Skip to content

Commit ec0b8ce

Browse files
committed
Merge branch 'release/v0.2.0'
2 parents 003932a + 51bc947 commit ec0b8ce

22 files changed

Lines changed: 1440 additions & 1141 deletions

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,3 +2,4 @@ venv
22
*.egg*
33
__pycache__
44
examples/data/
5+
*.svg

README.md

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,8 @@
1-
# Pytorch SOOM
1+
# Torch Numerical Optimization
22

3-
Implementation of second order optimization methods for Neural Networks.
3+
Implementation of numerical optimization methods for Neural Networks.
44

5-
Due to computational constraints, these methods are to be used with small Neural Networks as they require $O(p^3)$ space for a network with $p$ parameters.
5+
Due to computational constraints, methods like Newton-Raphson or Levenberg-Marquardt are to be used with small Neural Networks as they require $O(p^3)$ space for a network with $p$ parameters.
66

77
## References
88
[relevant paper](https://iopscience.iop.org/article/10.1088/1757-899X/495/1/012003/pdf)
@@ -11,8 +11,9 @@ Due to computational constraints, these methods are to be used with small Neural
1111

1212
- [x] Newton-Raphson
1313
- [x] Gauss-Newton
14-
- [x] Levemberg-Marquard (LM)
14+
- [x] Levenberg-Marquard (LM)
1515
- [x] Approximate Greatest Descent (AGD)
1616
- [ ] Conjugate Gradient
1717
- [ ] Quasi-Newton (LBFGS already in pytorch)
1818
- [ ] Hessian-free / truncated Newton
19+
- [ ] Stochastic Gradient Descent with Line Search

examples/quick_tests_agd.ipynb

Lines changed: 240 additions & 222 deletions
Large diffs are not rendered by default.

examples/quick_tests_baseline.ipynb

Lines changed: 220 additions & 221 deletions
Large diffs are not rendered by default.

examples/quick_tests_gauss_newton.ipynb

Lines changed: 129 additions & 119 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
"import torch\n",
1010
"import torch.nn as nn\n",
1111
"import torch.optim as optim\n",
12-
"import pytorch_soom\n",
12+
"import torch_numopt\n",
1313
"import numpy as np\n",
1414
"import matplotlib.pyplot as plt\n",
1515
"from torch.utils.data import DataLoader, TensorDataset\n",
@@ -23,7 +23,7 @@
2323
"metadata": {},
2424
"outputs": [],
2525
"source": [
26-
"device = 'cpu'"
26+
"device = \"cpu\""
2727
]
2828
},
2929
{
@@ -33,7 +33,7 @@
3333
"outputs": [],
3434
"source": [
3535
"class Net(nn.Module):\n",
36-
" def __init__(self, input_size, device='cpu'):\n",
36+
" def __init__(self, input_size, device=\"cpu\"):\n",
3737
" super().__init__()\n",
3838
" self.f1 = nn.Linear(input_size, 10, device=device)\n",
3939
" self.f2 = nn.Linear(10, 20, device=device)\n",
@@ -50,19 +50,27 @@
5050
" x = self.activation(self.f3(x))\n",
5151
" x = self.activation(self.f4(x))\n",
5252
" x = self.f5(x)\n",
53-
" \n",
54-
" return x\n"
53+
"\n",
54+
" return x"
5555
]
5656
},
5757
{
5858
"cell_type": "code",
5959
"execution_count": 4,
6060
"metadata": {},
61-
"outputs": [],
61+
"outputs": [
62+
{
63+
"name": "stdout",
64+
"output_type": "stream",
65+
"text": [
66+
"(442, 10)\n"
67+
]
68+
}
69+
],
6270
"source": [
63-
"X, y = load_diabetes(return_X_y = True, scaled=False)\n",
71+
"X, y = load_diabetes(return_X_y=True, scaled=False)\n",
6472
"# X, y = make_regression(n_samples=1000, n_features=100)\n",
65-
"# print(X.shape)\n",
73+
"print(X.shape)\n",
6674
"\n",
6775
"X_scaler = MinMaxScaler()\n",
6876
"X = X_scaler.fit_transform(X)\n",
@@ -76,128 +84,130 @@
7684
},
7785
{
7886
"cell_type": "code",
79-
"execution_count": 6,
87+
"execution_count": 5,
8088
"metadata": {},
8189
"outputs": [
8290
{
8391
"name": "stdout",
8492
"output_type": "stream",
8593
"text": [
86-
"epoch: 0, loss: 0.472525417804718\n",
87-
"epoch: 1, loss: 0.3582996428012848\n",
88-
"epoch: 2, loss: 0.21486149728298187\n",
89-
"epoch: 3, loss: 0.05209813266992569\n",
90-
"epoch: 4, loss: 0.04992499202489853\n",
91-
"epoch: 5, loss: 0.04917919635772705\n",
92-
"epoch: 6, loss: 0.04885313659906387\n",
93-
"epoch: 7, loss: 0.04869797080755234\n",
94-
"epoch: 8, loss: 0.04817237704992294\n",
95-
"epoch: 9, loss: 0.047917649149894714\n",
96-
"epoch: 10, loss: 0.04777839779853821\n",
97-
"epoch: 11, loss: 0.04744839668273926\n",
98-
"epoch: 12, loss: 0.04713824763894081\n",
99-
"epoch: 13, loss: 0.046973273158073425\n",
100-
"epoch: 14, loss: 0.04657743126153946\n",
101-
"epoch: 15, loss: 0.04626259580254555\n",
102-
"epoch: 16, loss: 0.04609677195549011\n",
103-
"epoch: 17, loss: 0.04590216651558876\n",
104-
"epoch: 18, loss: 0.04551844671368599\n",
105-
"epoch: 19, loss: 0.045348066836595535\n",
106-
"epoch: 20, loss: 0.045282550156116486\n",
107-
"epoch: 21, loss: 0.04486611485481262\n",
108-
"epoch: 22, loss: 0.04467124864459038\n",
109-
"epoch: 23, loss: 0.044562213122844696\n",
110-
"epoch: 24, loss: 0.04423132166266441\n",
111-
"epoch: 25, loss: 0.043990567326545715\n",
112-
"epoch: 26, loss: 0.043874021619558334\n",
113-
"epoch: 27, loss: 0.043532807379961014\n",
114-
"epoch: 28, loss: 0.043306369334459305\n",
115-
"epoch: 29, loss: 0.04319066181778908\n",
116-
"epoch: 30, loss: 0.042983170598745346\n",
117-
"epoch: 31, loss: 0.042671386152505875\n",
118-
"epoch: 32, loss: 0.04253864660859108\n",
119-
"epoch: 33, loss: 0.042385730892419815\n",
120-
"epoch: 34, loss: 0.0420558862388134\n",
121-
"epoch: 35, loss: 0.04190899059176445\n",
122-
"epoch: 36, loss: 0.04189492017030716\n",
123-
"epoch: 37, loss: 0.04145083203911781\n",
124-
"epoch: 38, loss: 0.04128744453191757\n",
125-
"epoch: 39, loss: 0.04127586632966995\n",
126-
"epoch: 40, loss: 0.04085730388760567\n",
127-
"epoch: 41, loss: 0.04069014638662338\n",
128-
"epoch: 42, loss: 0.04060281440615654\n",
129-
"epoch: 43, loss: 0.04029551148414612\n",
130-
"epoch: 44, loss: 0.04011178016662598\n",
131-
"epoch: 45, loss: 0.04002634435892105\n",
132-
"epoch: 46, loss: 0.03975091502070427\n",
133-
"epoch: 47, loss: 0.03956512361764908\n",
134-
"epoch: 48, loss: 0.039477184414863586\n",
135-
"epoch: 49, loss: 0.03926690295338631\n",
136-
"epoch: 50, loss: 0.03904862329363823\n",
137-
"epoch: 51, loss: 0.03895776346325874\n",
138-
"epoch: 52, loss: 0.03878797963261604\n",
139-
"epoch: 53, loss: 0.03855712711811066\n",
140-
"epoch: 54, loss: 0.0384608656167984\n",
141-
"epoch: 55, loss: 0.03836053982377052\n",
142-
"epoch: 56, loss: 0.03808850049972534\n",
143-
"epoch: 57, loss: 0.037987515330314636\n",
144-
"epoch: 58, loss: 0.0378975085914135\n",
145-
"epoch: 59, loss: 0.03763682767748833\n",
146-
"epoch: 60, loss: 0.03753144294023514\n",
147-
"epoch: 61, loss: 0.0374930240213871\n",
148-
"epoch: 62, loss: 0.03719552233815193\n",
149-
"epoch: 63, loss: 0.03709506243467331\n",
150-
"epoch: 64, loss: 0.03708367049694061\n",
151-
"epoch: 65, loss: 0.03679373860359192\n",
152-
"epoch: 66, loss: 0.03668645769357681\n",
153-
"epoch: 67, loss: 0.036631885915994644\n",
154-
"epoch: 68, loss: 0.03645823895931244\n",
155-
"epoch: 69, loss: 0.036306194961071014\n",
156-
"epoch: 70, loss: 0.03624321147799492\n",
157-
"epoch: 71, loss: 0.036088865250349045\n",
158-
"epoch: 72, loss: 0.03593679144978523\n",
159-
"epoch: 73, loss: 0.03587287291884422\n",
160-
"epoch: 74, loss: 0.035797119140625\n",
161-
"epoch: 75, loss: 0.03559079021215439\n",
162-
"epoch: 76, loss: 0.035518042743206024\n",
163-
"epoch: 77, loss: 0.03542740270495415\n",
164-
"epoch: 78, loss: 0.035248108208179474\n",
165-
"epoch: 79, loss: 0.03518311306834221\n",
166-
"epoch: 80, loss: 0.03511938080191612\n",
167-
"epoch: 81, loss: 0.0349227711558342\n",
168-
"epoch: 82, loss: 0.03486097231507301\n",
169-
"epoch: 83, loss: 0.03476268798112869\n",
170-
"epoch: 84, loss: 0.03461017087101936\n",
171-
"epoch: 85, loss: 0.03455572947859764\n",
172-
"epoch: 86, loss: 0.034471094608306885\n",
173-
"epoch: 87, loss: 0.03431500121951103\n",
174-
"epoch: 88, loss: 0.03426392003893852\n",
175-
"epoch: 89, loss: 0.03417592495679855\n",
176-
"epoch: 90, loss: 0.03403814882040024\n",
177-
"epoch: 91, loss: 0.03398758918046951\n",
178-
"epoch: 92, loss: 0.033923953771591187\n",
179-
"epoch: 93, loss: 0.03377244621515274\n",
180-
"epoch: 94, loss: 0.03372503072023392\n",
181-
"epoch: 95, loss: 0.03364565595984459\n",
182-
"epoch: 96, loss: 0.03352084010839462\n",
183-
"epoch: 97, loss: 0.03347627446055412\n",
184-
"epoch: 98, loss: 0.03339993581175804\n",
185-
"epoch: 99, loss: 0.033276401460170746\n"
94+
"epoch: 0, loss: 0.4598602056503296\n",
95+
"epoch: 1, loss: 0.459220290184021\n",
96+
"epoch: 2, loss: 0.4265614449977875\n",
97+
"epoch: 3, loss: 0.40943145751953125\n",
98+
"epoch: 4, loss: 0.3929181694984436\n",
99+
"epoch: 5, loss: 0.3746211528778076\n",
100+
"epoch: 6, loss: 0.3568786680698395\n",
101+
"epoch: 7, loss: 0.3395313322544098\n",
102+
"epoch: 8, loss: 0.3224108815193176\n",
103+
"epoch: 9, loss: 0.30596664547920227\n",
104+
"epoch: 10, loss: 0.2890184223651886\n",
105+
"epoch: 11, loss: 0.27232205867767334\n",
106+
"epoch: 12, loss: 0.252121239900589\n",
107+
"epoch: 13, loss: 0.23553621768951416\n",
108+
"epoch: 14, loss: 0.22030159831047058\n",
109+
"epoch: 15, loss: 0.20613881945610046\n",
110+
"epoch: 16, loss: 0.19282163679599762\n",
111+
"epoch: 17, loss: 0.1815483421087265\n",
112+
"epoch: 18, loss: 0.16843508183956146\n",
113+
"epoch: 19, loss: 0.15587976574897766\n",
114+
"epoch: 20, loss: 0.14433977007865906\n",
115+
"epoch: 21, loss: 0.13441333174705505\n",
116+
"epoch: 22, loss: 0.13338324427604675\n",
117+
"epoch: 23, loss: 0.12291105091571808\n",
118+
"epoch: 24, loss: 0.11294158548116684\n",
119+
"epoch: 25, loss: 0.10494914650917053\n",
120+
"epoch: 26, loss: 0.09662390500307083\n",
121+
"epoch: 27, loss: 0.0899830088019371\n",
122+
"epoch: 28, loss: 0.08319162577390671\n",
123+
"epoch: 29, loss: 0.07599890977144241\n",
124+
"epoch: 30, loss: 0.07022660970687866\n",
125+
"epoch: 31, loss: 0.06437348574399948\n",
126+
"epoch: 32, loss: 0.06383267045021057\n",
127+
"epoch: 33, loss: 0.05796610191464424\n",
128+
"epoch: 34, loss: 0.05324292182922363\n",
129+
"epoch: 35, loss: 0.04937686771154404\n",
130+
"epoch: 36, loss: 0.04662841558456421\n",
131+
"epoch: 37, loss: 0.042426615953445435\n",
132+
"epoch: 38, loss: 0.039919886738061905\n",
133+
"epoch: 39, loss: 0.03758644685149193\n",
134+
"epoch: 40, loss: 0.03577081859111786\n",
135+
"epoch: 41, loss: 0.03429730609059334\n",
136+
"epoch: 42, loss: 0.03257429599761963\n",
137+
"epoch: 43, loss: 0.03131904453039169\n",
138+
"epoch: 44, loss: 0.029763733968138695\n",
139+
"epoch: 45, loss: 0.028603361919522285\n",
140+
"epoch: 46, loss: 0.02832748368382454\n",
141+
"epoch: 47, loss: 0.027592387050390244\n",
142+
"epoch: 48, loss: 0.027274420484900475\n",
143+
"epoch: 49, loss: 0.025557763874530792\n",
144+
"epoch: 50, loss: 0.025391239672899246\n",
145+
"epoch: 51, loss: 0.024014215916395187\n",
146+
"epoch: 52, loss: 0.02375749684870243\n",
147+
"epoch: 53, loss: 0.023664424195885658\n",
148+
"epoch: 54, loss: 0.023138387128710747\n",
149+
"epoch: 55, loss: 0.022885076701641083\n",
150+
"epoch: 56, loss: 0.022653499618172646\n",
151+
"epoch: 57, loss: 0.022489318624138832\n",
152+
"epoch: 58, loss: 0.021514035761356354\n",
153+
"epoch: 59, loss: 0.021260851994156837\n",
154+
"epoch: 60, loss: 0.02007436379790306\n",
155+
"epoch: 61, loss: 0.01984875090420246\n",
156+
"epoch: 62, loss: 0.01916597969830036\n",
157+
"epoch: 63, loss: 0.018954439088702202\n",
158+
"epoch: 64, loss: 0.018762042745947838\n",
159+
"epoch: 65, loss: 0.018582148477435112\n",
160+
"epoch: 66, loss: 0.018494345247745514\n",
161+
"epoch: 67, loss: 0.018258405849337578\n",
162+
"epoch: 68, loss: 0.018063809722661972\n",
163+
"epoch: 69, loss: 0.017862489446997643\n",
164+
"epoch: 70, loss: 0.01769265905022621\n",
165+
"epoch: 71, loss: 0.017532240599393845\n",
166+
"epoch: 72, loss: 0.0173875093460083\n",
167+
"epoch: 73, loss: 0.017245948314666748\n",
168+
"epoch: 74, loss: 0.017108041793107986\n",
169+
"epoch: 75, loss: 0.016938941553235054\n",
170+
"epoch: 76, loss: 0.016798121854662895\n",
171+
"epoch: 77, loss: 0.016725268214941025\n",
172+
"epoch: 78, loss: 0.016584692522883415\n",
173+
"epoch: 79, loss: 0.016525041311979294\n",
174+
"epoch: 80, loss: 0.01648643985390663\n",
175+
"epoch: 81, loss: 0.016379758715629578\n",
176+
"epoch: 82, loss: 0.01636452041566372\n",
177+
"epoch: 83, loss: 0.01628103479743004\n",
178+
"epoch: 84, loss: 0.016180608421564102\n",
179+
"epoch: 85, loss: 0.016170047223567963\n",
180+
"epoch: 86, loss: 0.016050659120082855\n",
181+
"epoch: 87, loss: 0.015961414203047752\n",
182+
"epoch: 88, loss: 0.01591377705335617\n",
183+
"epoch: 89, loss: 0.01585211418569088\n",
184+
"epoch: 90, loss: 0.015741195529699326\n",
185+
"epoch: 91, loss: 0.01570090651512146\n",
186+
"epoch: 92, loss: 0.015616626478731632\n",
187+
"epoch: 93, loss: 0.015601896680891514\n",
188+
"epoch: 94, loss: 0.015534140169620514\n",
189+
"epoch: 95, loss: 0.015430513769388199\n",
190+
"epoch: 96, loss: 0.01536334864795208\n",
191+
"epoch: 97, loss: 0.015257438644766808\n",
192+
"epoch: 98, loss: 0.015197168104350567\n",
193+
"epoch: 99, loss: 0.015128728933632374\n"
186194
]
187195
}
188196
],
189197
"source": [
190-
"model = Net(input_size = X.shape[1], device=device)\n",
198+
"model = Net(input_size=X.shape[1], device=device)\n",
191199
"loss_fn = nn.MSELoss()\n",
192-
"opt = pytorch_soom.GaussNewton(model.parameters(), lr=1, model=model, c1=1e-4, tau=0.1, line_search_method='backtrack', line_search_cond='armijo')\n",
193-
"# opt = pytorch_soom.GaussNewton(model.parameters(), lr=1, model=model, c1=1e-4, tau=0.5, line_search_method='backtrack', line_search_cond='wolfe')\n",
194-
"# opt = pytorch_soom.GaussNewton(model.parameters(), lr=1, model=model, hessian_approx=False, c1=1e-4, tau=0.5, line_search_method='backtrack', line_search_cond='strong-wolfe')\n",
195-
"# opt = pytorch_soom.GaussNewton(model.parameters(), lr=1, model=model, hessian_approx=False, c1=1e-4, tau=0.5, line_search_method='backtrack', line_search_cond='goldstein')\n",
200+
"# loss_fn = nn.L1Loss()\n",
201+
"# loss_fn = nn.NLLLoss(reduction='mean')\n",
202+
"opt = torch_numopt.GaussNewton(model.parameters(), lr=1, model=model, c1=1e-4, tau=0.1, line_search_method=\"backtrack\", line_search_cond=\"armijo\")\n",
203+
"# opt = torch_numopt.GaussNewton(model.parameters(), lr=1, model=model, c1=1e-4, tau=0.5, line_search_method='backtrack', line_search_cond='wolfe')\n",
204+
"# opt = torch_numopt.GaussNewton(model.parameters(), lr=1, model=model, hessian_approx=False, c1=1e-4, tau=0.5, line_search_method='backtrack', line_search_cond='strong-wolfe')\n",
205+
"# opt = torch_numopt.GaussNewton(model.parameters(), lr=1, model=model, hessian_approx=False, c1=1e-4, tau=0.5, line_search_method='backtrack', line_search_cond='goldstein')\n",
196206
"\n",
197207
"all_loss = {}\n",
198208
"for epoch in range(100):\n",
199-
" print('epoch: ', epoch, end='')\n",
200-
" all_loss[epoch+1] = 0\n",
209+
" print(\"epoch: \", epoch, end=\"\")\n",
210+
" all_loss[epoch + 1] = 0\n",
201211
" for batch_idx, (b_x, b_y) in enumerate(data_loader):\n",
202212
" pre = model(b_x)\n",
203213
" loss = loss_fn(pre, b_y)\n",
@@ -207,9 +217,9 @@
207217
" # parameter update step based on optimizer\n",
208218
" opt.step(b_x, b_y, loss_fn)\n",
209219
"\n",
210-
" all_loss[epoch+1] += loss\n",
211-
" all_loss[epoch+1] /= len(data_loader)\n",
212-
" print(', loss: {}'.format(all_loss[epoch+1].detach().numpy().item()))"
220+
" all_loss[epoch + 1] += loss\n",
221+
" all_loss[epoch + 1] /= len(data_loader)\n",
222+
" print(\", loss: {}\".format(all_loss[epoch + 1].detach().numpy().item()))"
213223
]
214224
}
215225
],

0 commit comments

Comments
 (0)