|
@@ -21,6 +21,26 @@ type TestMetrics = Tuple[float, float] # (test_loss, test_acc)
|
|
|
type KLLossFn = Callable[[nn.Module], torch.Tensor | None]
|
|
type KLLossFn = Callable[[nn.Module], torch.Tensor | None]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
+def _model_device(model: nn.Module) -> torch.device:
|
|
|
|
|
+ first_param = next(model.parameters(), None)
|
|
|
|
|
+ if first_param is None:
|
|
|
|
|
+ return torch.device("cpu")
|
|
|
|
|
+ return first_param.device
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _move_batch_to_model_device(
|
|
|
|
|
+ model: nn.Module,
|
|
|
|
|
+ mri: torch.Tensor,
|
|
|
|
|
+ xls: torch.Tensor,
|
|
|
|
|
+ targets: torch.Tensor,
|
|
|
|
|
+) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
|
|
|
+ device = _model_device(model)
|
|
|
|
|
+ mri = mri.to(device, non_blocking=True)
|
|
|
|
|
+ xls = xls.to(device, non_blocking=True)
|
|
|
|
|
+ targets = targets.to(device, non_blocking=True)
|
|
|
|
|
+ return mri, xls, targets
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
def _batch_correct_and_total(
|
|
def _batch_correct_and_total(
|
|
|
outputs: torch.Tensor,
|
|
outputs: torch.Tensor,
|
|
|
targets: torch.Tensor,
|
|
targets: torch.Tensor,
|
|
@@ -84,6 +104,7 @@ def test_model(
|
|
|
test_progress.update(total=len(test_loader), advance=0)
|
|
test_progress.update(total=len(test_loader), advance=0)
|
|
|
for _, (mri, xls, targets, _) in enumerate(test_loader):
|
|
for _, (mri, xls, targets, _) in enumerate(test_loader):
|
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ mri, xls, targets = _move_batch_to_model_device(model, mri, xls, targets)
|
|
|
outputs = model((mri, xls))
|
|
outputs = model((mri, xls))
|
|
|
loss = criterion(outputs, targets)
|
|
loss = criterion(outputs, targets)
|
|
|
batch_size = mri.size(0)
|
|
batch_size = mri.size(0)
|
|
@@ -138,6 +159,7 @@ def train_epoch(
|
|
|
train_progress.update(total=len(train_loader), advance=0)
|
|
train_progress.update(total=len(train_loader), advance=0)
|
|
|
for _, (mri, xls, targets, _) in enumerate(train_loader):
|
|
for _, (mri, xls, targets, _) in enumerate(train_loader):
|
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ mri, xls, targets = _move_batch_to_model_device(model, mri, xls, targets)
|
|
|
optimizer.zero_grad()
|
|
optimizer.zero_grad()
|
|
|
outputs = model((mri, xls))
|
|
outputs = model((mri, xls))
|
|
|
loss = criterion(outputs, targets)
|
|
loss = criterion(outputs, targets)
|
|
@@ -165,6 +187,7 @@ def train_epoch(
|
|
|
with torch.no_grad():
|
|
with torch.no_grad():
|
|
|
for _, (mri, xls, targets, _) in enumerate(val_loader):
|
|
for _, (mri, xls, targets, _) in enumerate(val_loader):
|
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ mri, xls, targets = _move_batch_to_model_device(model, mri, xls, targets)
|
|
|
outputs = model((mri, xls))
|
|
outputs = model((mri, xls))
|
|
|
loss = criterion(outputs, targets)
|
|
loss = criterion(outputs, targets)
|
|
|
batch_size = mri.size(0)
|
|
batch_size = mri.size(0)
|
|
@@ -292,6 +315,7 @@ def test_model_bayesian(
|
|
|
with torch.no_grad():
|
|
with torch.no_grad():
|
|
|
for _, (mri, xls, targets, _) in enumerate(test_loader):
|
|
for _, (mri, xls, targets, _) in enumerate(test_loader):
|
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ mri, xls, targets = _move_batch_to_model_device(model, mri, xls, targets)
|
|
|
outputs = model((mri, xls))
|
|
outputs = model((mri, xls))
|
|
|
data_loss = cast(torch.Tensor, criterion(outputs, targets))
|
|
data_loss = cast(torch.Tensor, criterion(outputs, targets))
|
|
|
batch_size = mri.size(0)
|
|
batch_size = mri.size(0)
|
|
@@ -343,6 +367,7 @@ def train_epoch_bayesian(
|
|
|
train_progress.update(total=len(train_loader), advance=0)
|
|
train_progress.update(total=len(train_loader), advance=0)
|
|
|
for _, (mri, xls, targets, _) in enumerate(train_loader):
|
|
for _, (mri, xls, targets, _) in enumerate(train_loader):
|
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ mri, xls, targets = _move_batch_to_model_device(model, mri, xls, targets)
|
|
|
optimizer.zero_grad()
|
|
optimizer.zero_grad()
|
|
|
outputs = model((mri, xls))
|
|
outputs = model((mri, xls))
|
|
|
data_loss = cast(torch.Tensor, criterion(outputs, targets))
|
|
data_loss = cast(torch.Tensor, criterion(outputs, targets))
|
|
@@ -377,6 +402,7 @@ def train_epoch_bayesian(
|
|
|
with torch.no_grad():
|
|
with torch.no_grad():
|
|
|
for _, (mri, xls, targets, _) in enumerate(val_loader):
|
|
for _, (mri, xls, targets, _) in enumerate(val_loader):
|
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
_check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ mri, xls, targets = _move_batch_to_model_device(model, mri, xls, targets)
|
|
|
outputs = model((mri, xls))
|
|
outputs = model((mri, xls))
|
|
|
data_loss = cast(torch.Tensor, criterion(outputs, targets))
|
|
data_loss = cast(torch.Tensor, criterion(outputs, targets))
|
|
|
batch_size = mri.size(0)
|
|
batch_size = mri.size(0)
|