|
24 | 24 | "outputs": [], |
25 | 25 | "source": [ |
26 | 26 | "# 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", |
30 | 30 | "\n", |
31 | 31 | "# syft absolute\n", |
32 | 32 | "import syft as sy\n", |
|
43 | 43 | }, |
44 | 44 | "outputs": [], |
45 | 45 | "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)" |
47 | 47 | ] |
48 | 48 | }, |
49 | 49 | { |
|
67 | 67 | }, |
68 | 68 | "outputs": [], |
69 | 69 | "source": [ |
70 | | - "key = random.PRNGKey(42)" |
| 70 | + "# Set the random seed for reproducibility\n", |
| 71 | + "torch.manual_seed(42)" |
71 | 72 | ] |
72 | 73 | }, |
73 | 74 | { |
|
79 | 80 | }, |
80 | 81 | "outputs": [], |
81 | 82 | "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" |
83 | 86 | ] |
84 | 87 | }, |
85 | 88 | { |
86 | 89 | "cell_type": "code", |
87 | 90 | "execution_count": null, |
88 | 91 | "id": "6", |
89 | | - "metadata": { |
90 | | - "tags": [] |
91 | | - }, |
| 92 | + "metadata": {}, |
92 | 93 | "outputs": [], |
93 | 94 | "source": [ |
94 | | - "assert round(train_data.sum()) == 1602" |
| 95 | + "assert torch.round(train_data.sum()) == 1557" |
95 | 96 | ] |
96 | 97 | }, |
97 | 98 | { |
|
127 | 128 | }, |
128 | 129 | "outputs": [], |
129 | 130 | "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)" |
131 | 133 | ] |
132 | 134 | }, |
133 | 135 | { |
134 | 136 | "cell_type": "code", |
135 | 137 | "execution_count": null, |
136 | 138 | "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": {}, |
140 | 150 | "outputs": [], |
141 | 151 | "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", |
145 | 155 | " self.out_dims = out_dims\n", |
| 156 | + " self.linear1 = nn.Linear(784, 128)\n", |
| 157 | + " self.linear2 = nn.Linear(128, out_dims)\n", |
146 | 158 | "\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", |
152 | 164 | " return x\n", |
153 | 165 | "\n", |
154 | 166 | "\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" |
161 | 169 | ] |
162 | 170 | }, |
163 | 171 | { |
164 | 172 | "cell_type": "code", |
165 | 173 | "execution_count": null, |
166 | | - "id": "11", |
167 | | - "metadata": { |
168 | | - "tags": [] |
169 | | - }, |
| 174 | + "id": "12", |
| 175 | + "metadata": {}, |
170 | 176 | "outputs": [], |
171 | 177 | "source": [ |
172 | | - "weights = model.init(key, train.syft_action_data)" |
| 178 | + "weights = model.state_dict()" |
173 | 179 | ] |
174 | 180 | }, |
175 | 181 | { |
176 | 182 | "cell_type": "code", |
177 | 183 | "execution_count": null, |
178 | | - "id": "12", |
| 184 | + "id": "13", |
179 | 185 | "metadata": { |
180 | 186 | "tags": [] |
181 | 187 | }, |
|
187 | 193 | { |
188 | 194 | "cell_type": "code", |
189 | 195 | "execution_count": null, |
190 | | - "id": "13", |
| 196 | + "id": "14", |
191 | 197 | "metadata": { |
192 | 198 | "tags": [] |
193 | 199 | }, |
|
199 | 205 | { |
200 | 206 | "cell_type": "code", |
201 | 207 | "execution_count": null, |
202 | | - "id": "14", |
| 208 | + "id": "15", |
203 | 209 | "metadata": { |
204 | 210 | "tags": [] |
205 | 211 | }, |
|
211 | 217 | { |
212 | 218 | "cell_type": "code", |
213 | 219 | "execution_count": null, |
214 | | - "id": "15", |
| 220 | + "id": "16", |
215 | 221 | "metadata": { |
216 | 222 | "tags": [] |
217 | 223 | }, |
|
223 | 229 | { |
224 | 230 | "cell_type": "code", |
225 | 231 | "execution_count": null, |
226 | | - "id": "16", |
| 232 | + "id": "17", |
227 | 233 | "metadata": { |
228 | 234 | "tags": [] |
229 | 235 | }, |
|
235 | 241 | ")\n", |
236 | 242 | "def train_mlp(weights, data):\n", |
237 | 243 | " # 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", |
240 | 247 | "\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", |
244 | 251 | " self.out_dims = out_dims\n", |
| 252 | + " self.linear1 = nn.Linear(784, 128)\n", |
| 253 | + " self.linear2 = nn.Linear(128, out_dims)\n", |
245 | 254 | "\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", |
251 | 260 | " return x\n", |
252 | 261 | "\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", |
256 | 272 | "\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", |
260 | 273 | " return output" |
261 | 274 | ] |
262 | 275 | }, |
263 | 276 | { |
264 | 277 | "cell_type": "code", |
265 | 278 | "execution_count": null, |
266 | | - "id": "17", |
| 279 | + "id": "18", |
267 | 280 | "metadata": { |
268 | 281 | "tags": [] |
269 | 282 | }, |
|
276 | 289 | { |
277 | 290 | "cell_type": "code", |
278 | 291 | "execution_count": null, |
279 | | - "id": "18", |
280 | | - "metadata": { |
281 | | - "tags": [] |
282 | | - }, |
| 292 | + "id": "19", |
| 293 | + "metadata": {}, |
283 | 294 | "outputs": [], |
284 | 295 | "source": [ |
285 | | - "assert round(output.sum(), 2) == -0.86" |
| 296 | + "assert torch.allclose(torch.sum(output), torch.tensor(1.3907))" |
286 | 297 | ] |
287 | 298 | }, |
288 | 299 | { |
289 | 300 | "cell_type": "code", |
290 | 301 | "execution_count": null, |
291 | | - "id": "19", |
| 302 | + "id": "20", |
292 | 303 | "metadata": { |
293 | 304 | "tags": [] |
294 | 305 | }, |
|
301 | 312 | { |
302 | 313 | "cell_type": "code", |
303 | 314 | "execution_count": null, |
304 | | - "id": "20", |
| 315 | + "id": "21", |
305 | 316 | "metadata": { |
306 | 317 | "tags": [] |
307 | 318 | }, |
|
313 | 324 | { |
314 | 325 | "cell_type": "code", |
315 | 326 | "execution_count": null, |
316 | | - "id": "21", |
| 327 | + "id": "22", |
317 | 328 | "metadata": { |
318 | 329 | "tags": [] |
319 | 330 | }, |
|
326 | 337 | { |
327 | 338 | "cell_type": "code", |
328 | 339 | "execution_count": null, |
329 | | - "id": "22", |
| 340 | + "id": "23", |
330 | 341 | "metadata": { |
331 | 342 | "tags": [] |
332 | 343 | }, |
|
338 | 349 | { |
339 | 350 | "cell_type": "code", |
340 | 351 | "execution_count": null, |
341 | | - "id": "23", |
| 352 | + "id": "24", |
342 | 353 | "metadata": {}, |
343 | 354 | "outputs": [], |
344 | 355 | "source": [ |
|
348 | 359 | { |
349 | 360 | "cell_type": "code", |
350 | 361 | "execution_count": null, |
351 | | - "id": "24", |
| 362 | + "id": "25", |
352 | 363 | "metadata": { |
353 | 364 | "tags": [] |
354 | 365 | }, |
355 | 366 | "outputs": [], |
356 | 367 | "source": [ |
357 | | - "assert round(float(result.sum()), 2) == -0.86" |
| 368 | + "assert torch.allclose(torch.sum(result), torch.tensor(1.3907))" |
358 | 369 | ] |
359 | 370 | }, |
360 | 371 | { |
361 | 372 | "cell_type": "code", |
362 | 373 | "execution_count": null, |
363 | | - "id": "25", |
| 374 | + "id": "26", |
364 | 375 | "metadata": { |
365 | 376 | "tags": [] |
366 | 377 | }, |
|
373 | 384 | { |
374 | 385 | "cell_type": "code", |
375 | 386 | "execution_count": null, |
376 | | - "id": "26", |
| 387 | + "id": "27", |
377 | 388 | "metadata": {}, |
378 | 389 | "outputs": [], |
379 | 390 | "source": [] |
|
395 | 406 | "name": "python", |
396 | 407 | "nbconvert_exporter": "python", |
397 | 408 | "pygments_lexer": "ipython3", |
398 | | - "version": "3.11.4" |
| 409 | + "version": "3.12.2" |
399 | 410 | }, |
400 | 411 | "toc": { |
401 | 412 | "base_numbering": 1, |
|
0 commit comments