update with orthobridge

This commit is contained in:
rpotter6298
2026-05-15 10:27:24 +02:00
parent 4dd2dbc734
commit 32a801a572
20 changed files with 591 additions and 6 deletions
+4
View File
@@ -241,6 +241,8 @@ def run(
if not losses:
continue
loss = sum(losses) / len(losses)
if hasattr(bridge, "modify_loss"):
loss = bridge.modify_loss(loss)
opt.zero_grad(); loss.backward(); opt.step()
total_loss += loss.item() * len(y_t)
total_n += len(y_t)
@@ -254,6 +256,8 @@ def run(
if logits is None:
continue
loss = F.cross_entropy(logits, y_t, weight=cw)
if hasattr(bridge, "modify_loss"):
loss = bridge.modify_loss(loss)
opt.zero_grad(); loss.backward(); opt.step()
total_correct += int((logits.argmax(1) == y_t).sum())
total_loss += loss.item() * len(y_t)
+4
View File
@@ -309,6 +309,8 @@ def _parallel_fusion(
if not losses:
continue
loss = sum(losses) / len(losses)
if hasattr(bridge, "modify_loss"):
loss = bridge.modify_loss(loss)
ctx["opt"].zero_grad(); loss.backward(); ctx["opt"].step()
total_loss += loss.item() * len(y_t)
total_n += len(y_t)
@@ -322,6 +324,8 @@ def _parallel_fusion(
if logits is None:
continue
loss = F.cross_entropy(logits, y_t, weight=ctx["class_weights"])
if hasattr(bridge, "modify_loss"):
loss = bridge.modify_loss(loss)
ctx["opt"].zero_grad(); loss.backward(); ctx["opt"].step()
total_correct += int((logits.argmax(1) == y_t).sum())
total_loss += loss.item() * len(y_t)