SACFlow / sacflow /utils /misc.py
sathiiii's picture
Add SACFlow source code
d65ae7d verified
Raw
History Blame Contribute Delete
1.11 kB
from __future__ import annotations
import os
import random
from pathlib import Path
import numpy as np
import torch
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
os.environ["PYTHONHASHSEED"] = str(seed)
def ensure_dir(path: str | Path) -> Path:
p = Path(path)
p.mkdir(parents=True, exist_ok=True)
return p
def count_trainable(model: torch.nn.Module) -> tuple[int, int]:
total = sum(p.numel() for p in model.parameters())
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
return trainable, total
def move_to_device(batch, device):
if isinstance(batch, torch.Tensor):
return batch.to(device, non_blocking=True)
if isinstance(batch, dict):
return {k: move_to_device(v, device) for k, v in batch.items()}
if isinstance(batch, (list, tuple)):
return type(batch)(move_to_device(v, device) for v in batch)
return batch
def unwrap_model(model):
return model.module if hasattr(model, "module") else model