Skip to content

Commit 870456b

Browse files
authored
Merge pull request #8837 from khoaguin/remove-jax-haiku
Remove Jax and Haiku. Use Torch instead
2 parents cb6dfe6 + f41c592 commit 870456b

16 files changed

Lines changed: 339 additions & 286 deletions

‎README.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ domain_client = sy.login(
6161
- <a href="notebooks/api/0.8/01-submit-code.ipynb">01-submit-code.ipynb</a>
6262
- <a href="notebooks/api/0.8/02-review-code-and-approve.ipynb">02-review-code-and-approve.ipynb</a>
6363
- <a href="notebooks/api/0.8/03-data-scientist-download-result.ipynb">03-data-scientist-download-result.ipynb</a>
64-
- <a href="notebooks/api/0.8/04-jax-example.ipynb">04-jax-example.ipynb</a>
64+
- <a href="notebooks/api/0.8/04-pytorch-example.ipynb">04-pytorch-example.ipynb</a>
6565
- <a href="notebooks/api/0.8/05-custom-policy.ipynb">05-custom-policy.ipynb</a>
6666
- <a href="notebooks/api/0.8/06-multiple-code-requests.ipynb">06-multiple-code-requests.ipynb</a>
6767
- <a href="notebooks/api/0.8/07-domain-register-control-flow.ipynb">07-domain-register-control-flow.ipynb</a>

notebooks/api/0.8/04-jax-example.ipynb renamed to notebooks/api/0.8/04-pytorch-example.ipynb

Lines changed: 81 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -24,9 +24,9 @@
2424
"outputs": [],
2525
"source": [
2626
"# third party\n",
27-
"import haiku as hk\n",
28-
"import jax\n",
29-
"from jax import random\n",
27+
"import torch\n",
28+
"import torch.nn as nn\n",
29+
"import torch.nn.functional as F\n",
3030
"\n",
3131
"# syft absolute\n",
3232
"import syft as sy\n",
@@ -43,7 +43,7 @@
4343
},
4444
"outputs": [],
4545
"source": [
46-
"node = sy.orchestra.launch(name=\"test-domain-1\", port=\"auto\", dev_mode=True)"
46+
"node = sy.orchestra.launch(name=\"test-domain-1\", dev_mode=True, reset=True)"
4747
]
4848
},
4949
{
@@ -67,7 +67,8 @@
6767
},
6868
"outputs": [],
6969
"source": [
70-
"key = random.PRNGKey(42)"
70+
"# Set the random seed for reproducibility\n",
71+
"torch.manual_seed(42)"
7172
]
7273
},
7374
{
@@ -79,19 +80,19 @@
7980
},
8081
"outputs": [],
8182
"source": [
82-
"train_data = random.uniform(key, shape=(4, 28, 28, 1))"
83+
"# Generate random data\n",
84+
"train_data = torch.rand((4, 28, 28, 1))\n",
85+
"train_data.shape"
8386
]
8487
},
8588
{
8689
"cell_type": "code",
8790
"execution_count": null,
8891
"id": "6",
89-
"metadata": {
90-
"tags": []
91-
},
92+
"metadata": {},
9293
"outputs": [],
9394
"source": [
94-
"assert round(train_data.sum()) == 1602"
95+
"assert torch.round(train_data.sum()) == 1557"
9596
]
9697
},
9798
{
@@ -127,55 +128,60 @@
127128
},
128129
"outputs": [],
129130
"source": [
130-
"train_domain_obj = domain_client.api.services.action.set(train)"
131+
"train_domain_obj = domain_client.api.services.action.set(train)\n",
132+
"type(train_domain_obj)"
131133
]
132134
},
133135
{
134136
"cell_type": "code",
135137
"execution_count": null,
136138
"id": "10",
137-
"metadata": {
138-
"tags": []
139-
},
139+
"metadata": {},
140+
"outputs": [],
141+
"source": [
142+
"assert torch.round(train_domain_obj.syft_action_data.sum()) == 1557"
143+
]
144+
},
145+
{
146+
"cell_type": "code",
147+
"execution_count": null,
148+
"id": "11",
149+
"metadata": {},
140150
"outputs": [],
141151
"source": [
142-
"class MLP(hk.Module):\n",
143-
" def __init__(self, out_dims, name=None):\n",
144-
" super().__init__(name=name)\n",
152+
"class MLP(nn.Module):\n",
153+
" def __init__(self, out_dims):\n",
154+
" super().__init__()\n",
145155
" self.out_dims = out_dims\n",
156+
" self.linear1 = nn.Linear(784, 128)\n",
157+
" self.linear2 = nn.Linear(128, out_dims)\n",
146158
"\n",
147-
" def __call__(self, x):\n",
148-
" x = x.reshape((x.shape[0], -1))\n",
149-
" x = hk.Linear(128)(x)\n",
150-
" x = jax.nn.relu(x)\n",
151-
" x = hk.Linear(self.out_dims)(x)\n",
159+
" def forward(self, x):\n",
160+
" x = x.view(x.size(0), -1)\n",
161+
" x = self.linear1(x)\n",
162+
" x = F.relu(x)\n",
163+
" x = self.linear2(x)\n",
152164
" return x\n",
153165
"\n",
154166
"\n",
155-
"def _forward_fn_linear1(x):\n",
156-
" module = MLP(out_dims=10)\n",
157-
" return module(x)\n",
158-
"\n",
159-
"\n",
160-
"model = hk.transform(_forward_fn_linear1)"
167+
"model = MLP(out_dims=10)\n",
168+
"model"
161169
]
162170
},
163171
{
164172
"cell_type": "code",
165173
"execution_count": null,
166-
"id": "11",
167-
"metadata": {
168-
"tags": []
169-
},
174+
"id": "12",
175+
"metadata": {},
170176
"outputs": [],
171177
"source": [
172-
"weights = model.init(key, train.syft_action_data)"
178+
"weights = model.state_dict()"
173179
]
174180
},
175181
{
176182
"cell_type": "code",
177183
"execution_count": null,
178-
"id": "12",
184+
"id": "13",
179185
"metadata": {
180186
"tags": []
181187
},
@@ -187,7 +193,7 @@
187193
{
188194
"cell_type": "code",
189195
"execution_count": null,
190-
"id": "13",
196+
"id": "14",
191197
"metadata": {
192198
"tags": []
193199
},
@@ -199,7 +205,7 @@
199205
{
200206
"cell_type": "code",
201207
"execution_count": null,
202-
"id": "14",
208+
"id": "15",
203209
"metadata": {
204210
"tags": []
205211
},
@@ -211,7 +217,7 @@
211217
{
212218
"cell_type": "code",
213219
"execution_count": null,
214-
"id": "15",
220+
"id": "16",
215221
"metadata": {
216222
"tags": []
217223
},
@@ -223,7 +229,7 @@
223229
{
224230
"cell_type": "code",
225231
"execution_count": null,
226-
"id": "16",
232+
"id": "17",
227233
"metadata": {
228234
"tags": []
229235
},
@@ -235,35 +241,42 @@
235241
")\n",
236242
"def train_mlp(weights, data):\n",
237243
" # third party\n",
238-
" import haiku as hk\n",
239-
" import jax\n",
244+
" import torch\n",
245+
" import torch.nn as nn\n",
246+
" import torch.nn.functional as F\n",
240247
"\n",
241-
" class MLP(hk.Module):\n",
242-
" def __init__(self, out_dims, name=None):\n",
243-
" super().__init__(name=name)\n",
248+
" class MLP(nn.Module):\n",
249+
" def __init__(self, out_dims):\n",
250+
" super().__init__()\n",
244251
" self.out_dims = out_dims\n",
252+
" self.linear1 = nn.Linear(784, 128)\n",
253+
" self.linear2 = nn.Linear(128, out_dims)\n",
245254
"\n",
246-
" def __call__(self, x):\n",
247-
" x = x.reshape((x.shape[0], -1))\n",
248-
" x = hk.Linear(128)(x)\n",
249-
" x = jax.nn.relu(x)\n",
250-
" x = hk.Linear(self.out_dims)(x)\n",
255+
" def forward(self, x):\n",
256+
" x = x.view(x.size(0), -1)\n",
257+
" x = self.linear1(x)\n",
258+
" x = F.relu(x)\n",
259+
" x = self.linear2(x)\n",
251260
" return x\n",
252261
"\n",
253-
" def _forward_fn_linear1(x):\n",
254-
" module = MLP(out_dims=10)\n",
255-
" return module(x)\n",
262+
" # Initialize the model\n",
263+
" model = MLP(out_dims=10)\n",
264+
"\n",
265+
" # Load weights into the model\n",
266+
" model.load_state_dict(weights)\n",
267+
"\n",
268+
" # Perform a forward pass\n",
269+
" model.eval() # Set the model to evaluation mode\n",
270+
" with torch.no_grad(): # Disable gradient calculation\n",
271+
" output = model(data)\n",
256272
"\n",
257-
" model = hk.transform(_forward_fn_linear1)\n",
258-
" rng_key = jax.random.PRNGKey(42)\n",
259-
" output = model.apply(params=weights, x=data, rng=rng_key)\n",
260273
" return output"
261274
]
262275
},
263276
{
264277
"cell_type": "code",
265278
"execution_count": null,
266-
"id": "17",
279+
"id": "18",
267280
"metadata": {
268281
"tags": []
269282
},
@@ -276,19 +289,17 @@
276289
{
277290
"cell_type": "code",
278291
"execution_count": null,
279-
"id": "18",
280-
"metadata": {
281-
"tags": []
282-
},
292+
"id": "19",
293+
"metadata": {},
283294
"outputs": [],
284295
"source": [
285-
"assert round(output.sum(), 2) == -0.86"
296+
"assert torch.allclose(torch.sum(output), torch.tensor(1.3907))"
286297
]
287298
},
288299
{
289300
"cell_type": "code",
290301
"execution_count": null,
291-
"id": "19",
302+
"id": "20",
292303
"metadata": {
293304
"tags": []
294305
},
@@ -301,7 +312,7 @@
301312
{
302313
"cell_type": "code",
303314
"execution_count": null,
304-
"id": "20",
315+
"id": "21",
305316
"metadata": {
306317
"tags": []
307318
},
@@ -313,7 +324,7 @@
313324
{
314325
"cell_type": "code",
315326
"execution_count": null,
316-
"id": "21",
327+
"id": "22",
317328
"metadata": {
318329
"tags": []
319330
},
@@ -326,7 +337,7 @@
326337
{
327338
"cell_type": "code",
328339
"execution_count": null,
329-
"id": "22",
340+
"id": "23",
330341
"metadata": {
331342
"tags": []
332343
},
@@ -338,7 +349,7 @@
338349
{
339350
"cell_type": "code",
340351
"execution_count": null,
341-
"id": "23",
352+
"id": "24",
342353
"metadata": {},
343354
"outputs": [],
344355
"source": [
@@ -348,19 +359,19 @@
348359
{
349360
"cell_type": "code",
350361
"execution_count": null,
351-
"id": "24",
362+
"id": "25",
352363
"metadata": {
353364
"tags": []
354365
},
355366
"outputs": [],
356367
"source": [
357-
"assert round(float(result.sum()), 2) == -0.86"
368+
"assert torch.allclose(torch.sum(result), torch.tensor(1.3907))"
358369
]
359370
},
360371
{
361372
"cell_type": "code",
362373
"execution_count": null,
363-
"id": "25",
374+
"id": "26",
364375
"metadata": {
365376
"tags": []
366377
},
@@ -373,7 +384,7 @@
373384
{
374385
"cell_type": "code",
375386
"execution_count": null,
376-
"id": "26",
387+
"id": "27",
377388
"metadata": {},
378389
"outputs": [],
379390
"source": []
@@ -395,7 +406,7 @@
395406
"name": "python",
396407
"nbconvert_exporter": "python",
397408
"pygments_lexer": "ipython3",
398-
"version": "3.11.4"
409+
"version": "3.12.2"
399410
},
400411
"toc": {
401412
"base_numbering": 1,

‎notebooks/tutorials/deployments/01-deploy-python.ipynb‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -154,7 +154,7 @@
154154
"- [01-submit-code.ipynb](../../api/0.8/01-submit-code.ipynb)\n",
155155
"- [02-review-code-and-approve.ipynb](../../api/0.8/02-review-code-and-approve.ipynb)\n",
156156
"- [03-data-scientist-download-result.ipynb](../../api/0.8/03-data-scientist-download-result.ipynb)\n",
157-
"- [04-jax-example.ipynb](../../api/0.8/04-jax-example.ipynb)\n",
157+
"- [04-pytorch-example.ipynb](../../api/0.8/04-pytorch-example.ipynb)\n",
158158
"- [05-custom-policy.ipynb](../../api/0.8/05-custom-policy.ipynb)\n",
159159
"- [06-multiple-code-requests.ipynb](../../api/0.8/06-multiple-code-requests.ipynb)\n",
160160
"- [07-domain-register-control-flow.ipynb](../../api/0.8/07-domain-register-control-flow.ipynb)\n",

‎notebooks/tutorials/deployments/02-deploy-container.ipynb‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -148,7 +148,7 @@
148148
"- [01-submit-code.ipynb](../../api/0.8/01-submit-code.ipynb)\n",
149149
"- [02-review-code-and-approve.ipynb](../../api/0.8/02-review-code-and-approve.ipynb)\n",
150150
"- [03-data-scientist-download-result.ipynb](../../api/0.8/03-data-scientist-download-result.ipynb)\n",
151-
"- [04-jax-example.ipynb](../../api/0.8/04-jax-example.ipynb)\n",
151+
"- [04-pytorch-example.ipynb](../../api/0.8/04-pytorch-example.ipynb)\n",
152152
"- [05-custom-policy.ipynb](../../api/0.8/05-custom-policy.ipynb)\n",
153153
"- [06-multiple-code-requests.ipynb](../../api/0.8/06-multiple-code-requests.ipynb)\n",
154154
"- [07-domain-register-control-flow.ipynb](../../api/0.8/07-domain-register-control-flow.ipynb)\n",

‎notebooks/tutorials/deployments/03-deploy-k8s-k3d.ipynb‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -156,7 +156,7 @@
156156
"- [01-submit-code.ipynb](../../api/0.8/01-submit-code.ipynb)\n",
157157
"- [02-review-code-and-approve.ipynb](../../api/0.8/02-review-code-and-approve.ipynb)\n",
158158
"- [03-data-scientist-download-result.ipynb](../../api/0.8/03-data-scientist-download-result.ipynb)\n",
159-
"- [04-jax-example.ipynb](../../api/0.8/04-jax-example.ipynb)\n",
159+
"- [04-pytorch-example.ipynb](../../api/0.8/04-pytorch-example.ipynb)\n",
160160
"- [05-custom-policy.ipynb](../../api/0.8/05-custom-policy.ipynb)\n",
161161
"- [06-multiple-code-requests.ipynb](../../api/0.8/06-multiple-code-requests.ipynb)\n",
162162
"- [07-domain-register-control-flow.ipynb](../../api/0.8/07-domain-register-control-flow.ipynb)\n",
@@ -167,6 +167,11 @@
167167
"\n",
168168
"Feel free to explore these notebooks to get started with PySyft and unlock its full potential for privacy-preserving machine learning!"
169169
]
170+
},
171+
{
172+
"cell_type": "markdown",
173+
"metadata": {},
174+
"source": []
170175
}
171176
],
172177
"metadata": {

0 commit comments

Comments
 (0)