|
32 | 32 | "import numpy as np\n", |
33 | 33 | "from tqdm import tqdm\n", |
34 | 34 | "import matplotlib.pyplot as plt\n", |
35 | | - "from torch.utils.data import DataLoader\n", |
36 | 35 | "from collections import OrderedDict" |
37 | 36 | ] |
38 | 37 | }, |
|
429 | 428 | }, |
430 | 429 | { |
431 | 430 | "cell_type": "code", |
432 | | - "execution_count": 25, |
| 431 | + "execution_count": null, |
433 | 432 | "id": "70c2c2fa", |
434 | 433 | "metadata": {}, |
435 | 434 | "outputs": [], |
436 | 435 | "source": [ |
437 | 436 | "seed = 42\n", |
438 | | - "import sys\n", |
| 437 | + "import sys # noqa: E402\n", |
439 | 438 | "\n", |
440 | 439 | "sys.path.append(\"/srv/user/turishcheva/sensorium_replicate/sensorium_2023/\")\n", |
441 | 440 | "sys.path.append(\"/srv/user/turishcheva/sensorium_replicate/neuralpredictors/\")\n", |
442 | | - "import torch\n", |
443 | | - "from nnfabrik.utility.nn_helpers import set_random_seed\n", |
| 441 | + "import torch # noqa: E402\n", |
| 442 | + "from nnfabrik.utility.nn_helpers import set_random_seed # noqa: E402\n", |
444 | 443 | "\n", |
445 | 444 | "set_random_seed(seed)\n", |
446 | 445 | "\n", |
447 | | - "from sensorium.datasets.mouse_video_loaders import mouse_video_loader\n", |
448 | | - "from sensorium.utility.scores import get_correlations\n", |
449 | | - "from nnfabrik.builder import get_trainer\n", |
450 | | - "from sensorium.models.make_model import make_video_model" |
| 446 | + "from nnfabrik.builder import get_trainer # noqa: E402\n", |
| 447 | + "from sensorium.models.make_model import make_video_model # noqa: E402" |
451 | 448 | ] |
452 | 449 | }, |
453 | 450 | { |
454 | 451 | "cell_type": "code", |
455 | | - "execution_count": 26, |
| 452 | + "execution_count": null, |
456 | 453 | "id": "78705901", |
457 | 454 | "metadata": {}, |
458 | 455 | "outputs": [], |
459 | 456 | "source": [ |
460 | | - "factorised_3D_core_dict = dict(\n", |
461 | | - " input_channels=1, # increase if behaviour is used\n", |
462 | | - " hidden_channels=[32, 64, 128],\n", |
463 | | - " spatial_input_kernel=(11, 11),\n", |
464 | | - " temporal_input_kernel=11,\n", |
465 | | - " spatial_hidden_kernel=(5, 5),\n", |
466 | | - " temporal_hidden_kernel=5,\n", |
467 | | - " stride=1,\n", |
468 | | - " layers=3,\n", |
469 | | - " gamma_input_spatial=10,\n", |
470 | | - " gamma_input_temporal=0.01,\n", |
471 | | - " bias=True,\n", |
472 | | - " hidden_nonlinearities=\"elu\",\n", |
473 | | - " x_shift=0,\n", |
474 | | - " y_shift=0,\n", |
475 | | - " batch_norm=True,\n", |
476 | | - " laplace_padding=None,\n", |
477 | | - " input_regularizer=\"LaplaceL2norm\",\n", |
478 | | - " padding=False,\n", |
479 | | - " final_nonlin=True,\n", |
480 | | - " momentum=0.7,\n", |
481 | | - ")\n", |
| 457 | + "factorised_3D_core_dict = {\n", |
| 458 | + " \"input_channels\": 1, # increase if behaviour is used\n", |
| 459 | + " \"hidden_channels\": [32, 64, 128],\n", |
| 460 | + " \"spatial_input_kernel\": (11, 11),\n", |
| 461 | + " \"temporal_input_kernel\": 11,\n", |
| 462 | + " \"spatial_hidden_kernel\": (5, 5),\n", |
| 463 | + " \"temporal_hidden_kernel\": 5,\n", |
| 464 | + " \"stride\": 1,\n", |
| 465 | + " \"layers\": 3,\n", |
| 466 | + " \"gamma_input_spatial\": 10,\n", |
| 467 | + " \"gamma_input_temporal\": 0.01,\n", |
| 468 | + " \"bias\": True,\n", |
| 469 | + " \"hidden_nonlinearities\": \"elu\",\n", |
| 470 | + " \"x_shift\": 0,\n", |
| 471 | + " \"y_shift\": 0,\n", |
| 472 | + " \"batch_norm\": True,\n", |
| 473 | + " \"laplace_padding\": None,\n", |
| 474 | + " \"input_regularizer\": \"LaplaceL2norm\",\n", |
| 475 | + " \"padding\": False,\n", |
| 476 | + " \"final_nonlin\": True,\n", |
| 477 | + " \"momentum\": 0.7,\n", |
| 478 | + "}\n", |
482 | 479 | "\n", |
483 | 480 | "\n", |
484 | 481 | "shifter_dict = None\n", |
485 | 482 | "\n", |
486 | 483 | "\n", |
487 | | - "readout_dict = dict(\n", |
488 | | - " bias=True,\n", |
489 | | - " init_mu_range=0.2,\n", |
490 | | - " init_sigma=1.0,\n", |
491 | | - " gamma_readout=0.0,\n", |
492 | | - " gauss_type=\"full\",\n", |
493 | | - " # grid_mean_predictor=None,\n", |
494 | | - " grid_mean_predictor={\n", |
| 484 | + "readout_dict = {\n", |
| 485 | + " \"bias\": True,\n", |
| 486 | + " \"init_mu_range\": 0.2,\n", |
| 487 | + " \"init_sigma\": 1.0,\n", |
| 488 | + " \"gamma_readout\": 0.0,\n", |
| 489 | + " \"gauss_type\": \"full\",\n", |
| 490 | + " # grid_mean_predictor=None,\n", |
| 491 | + " \"grid_mean_predictor\": {\n", |
495 | 492 | " \"type\": \"cortex\",\n", |
496 | 493 | " \"input_dimensions\": 2,\n", |
497 | 494 | " \"hidden_layers\": 1,\n", |
498 | 495 | " \"hidden_features\": 30,\n", |
499 | 496 | " \"final_tanh\": True,\n", |
500 | 497 | " },\n", |
501 | | - " share_features=False,\n", |
502 | | - " share_grid=False,\n", |
503 | | - " shared_match_ids=None,\n", |
504 | | - " gamma_grid_dispersion=0.0,\n", |
505 | | - ")" |
| 498 | + " \"share_features\": False,\n", |
| 499 | + " \"share_grid\": False,\n", |
| 500 | + " \"shared_match_ids\": None,\n", |
| 501 | + " \"gamma_grid_dispersion\": 0.0,\n", |
| 502 | + "}" |
506 | 503 | ] |
507 | 504 | }, |
508 | 505 | { |
|
744 | 741 | "source": [ |
745 | 742 | "import yaml\n", |
746 | 743 | "import os\n", |
747 | | - "from tqdm import tqdm\n", |
748 | 744 | "from collections import Counter" |
749 | 745 | ] |
750 | 746 | }, |
|
763 | 759 | "metadata": {}, |
764 | 760 | "outputs": [], |
765 | 761 | "source": [ |
| 762 | + "import pickle\n", |
| 763 | + "\n", |
766 | 764 | "path_to_old = ...\n", |
767 | 765 | "path_to_new = ...\n", |
768 | 766 | "path_to_save_matching = ...\n", |
|
782 | 780 | " )\n", |
783 | 781 | "\n", |
784 | 782 | " for file in tqdm(os.listdir(yaml_pre_path)):\n", |
785 | | - " with open(f\"{yaml_pre_path}{file}\", \"r\") as f:\n", |
| 783 | + " with open(f\"{yaml_pre_path}{file}\") as f:\n", |
786 | 784 | " data = yaml.safe_load(f)\n", |
787 | 785 | " if data[\"modality\"] != \"blank\":\n", |
788 | 786 | " if len(np.where(trial_idx_prev == data[\"trial_idx\"])[0]) == 0:\n", |
|
817 | 815 | "metadata": {}, |
818 | 816 | "outputs": [], |
819 | 817 | "source": [ |
| 818 | + "import pickle\n", |
| 819 | + "\n", |
820 | 820 | "for m in [\n", |
821 | 821 | " \"dynamic29623-4-9-Video-full\",\n", |
822 | 822 | " \"dynamic29647-19-8-Video-full\",\n", |
|
845 | 845 | " trial_idx_prev = np.asarray(trial_idx_prev)\n", |
846 | 846 | "\n", |
847 | 847 | " for file in tqdm(os.listdir(yaml_pre_path)):\n", |
848 | | - " with open(f\"{yaml_pre_path}{file}\", \"r\") as f:\n", |
| 848 | + " with open(f\"{yaml_pre_path}{file}\") as f:\n", |
849 | 849 | " data = yaml.safe_load(f)\n", |
850 | 850 | " if data[\"modality\"] != \"blank\":\n", |
851 | 851 | " if len(np.where(trial_idx_prev == data[\"trial_idx\"])[0]) == 0:\n", |
|
0 commit comments