Skip to content

Commit b8d27f0

Browse files
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
1 parent 87f1c28 commit b8d27f0

6 files changed

Lines changed: 26 additions & 21 deletions

File tree

codes/benchmark/bench_fcts.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -460,7 +460,9 @@ def evaluate_iterative_predictions(
460460
# We predict steps 1..(chunk_len-1) relative to the provided init state (index 0).
461461
# Map these to global indices [start+1 .. end] inclusively.
462462
if i == 0:
463-
iterative_preds[:, start : end + 1, :] = preds_chunk[:, : model.n_timesteps, :].detach().cpu().numpy()
463+
iterative_preds[:, start : end + 1, :] = (
464+
preds_chunk[:, : model.n_timesteps, :].detach().cpu().numpy()
465+
)
464466
iterative_preds[:, start + 1 : end + 1, :] = (
465467
preds_chunk[:, 1 : model.n_timesteps, :].detach().cpu().numpy()
466468
)

codes/benchmark/bench_utils.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -261,7 +261,9 @@ def count_trainable_parameters(model: torch.nn.Module) -> int:
261261
return sum(p.numel() for p in model.parameters() if p.requires_grad)
262262

263263

264-
def measure_memory_footprint(model: torch.nn.Module, inputs: tuple, device: torch.device) -> dict:
264+
def measure_memory_footprint(
265+
model: torch.nn.Module, inputs: tuple, device: torch.device
266+
) -> dict:
265267
"""
266268
Measure peak GPU memory usage for forward/backward passes.
267269

codes/train/__init__.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,10 @@
11
from .train_fcts import (
2+
DummyLock,
3+
create_task_list_for_surrogate,
24
parallel_training,
35
sequential_training,
46
train_and_save_model,
5-
create_task_list_for_surrogate,
67
worker,
7-
DummyLock,
88
)
99

1010
__all__ = [

datasets/_data_analysis/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,5 +4,5 @@
44
"get_data_subset",
55
"create_dataset",
66
"normalize_data",
7-
"download_data"
7+
"download_data",
88
]

test/test_training_pipeline.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,20 +1,20 @@
11
import threading
22
from queue import Queue
3-
import pytest
43
from unittest.mock import Mock, patch
54

5+
import pytest
6+
67
# import *from* the module that actually defines them
78
from codes.train.train_fcts import (
89
DummyLock,
910
create_task_list_for_surrogate,
10-
train_and_save_model,
11-
worker,
1211
parallel_training,
1312
sequential_training,
13+
train_and_save_model,
14+
worker,
1415
)
1516
from codes.utils import load_task_list, save_task_list
1617

17-
1818
# — fixtures —
1919

2020

test/test_utils.py

Lines changed: 13 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,27 @@
11
import os
2-
import time
32
import random
4-
import yaml
5-
import torch
3+
import time
4+
65
import numpy as np
76
import pytest
7+
import torch
8+
import yaml
89

910
from codes.utils import (
10-
read_yaml_config,
11-
time_execution,
11+
batch_factor_to_float,
12+
check_training_status,
1213
create_model_dir,
14+
determine_batch_size,
1315
get_progress_bar,
1416
load_and_save_config,
15-
set_random_seeds,
16-
nice_print,
17+
load_task_list,
1718
make_description,
18-
worker_init_fn,
19+
nice_print,
20+
read_yaml_config,
1921
save_task_list,
20-
load_task_list,
21-
check_training_status,
22-
determine_batch_size,
23-
batch_factor_to_float,
22+
set_random_seeds,
23+
time_execution,
24+
worker_init_fn,
2425
)
2526

2627

0 commit comments

Comments
 (0)