| import torch as t |
| import argparse |
| from torch.utils.data import DataLoader, TensorDataset, random_split |
| from typing import List, Tuple, Dict, Any, Union, Callable, cast |
|
|
| def str_to_bool(value): |
| if isinstance(value, bool): |
| return value |
| if value.lower() in ('yes', 'true', 't', 'y', '1'): |
| return True |
| elif value.lower() in ('no', 'false', 'f', 'n', '0'): |
| return False |
| else: |
| raise argparse.ArgumentTypeError('Boolean value expected.') |
|
|
| def freeze_model(model: t.nn.Module) -> None: |
| for param in model.parameters(): |
| param.requires_grad = False |
| |
| def calculate_kl_div(inputs, targets, reduction='batchmean'): |
| |
| input_log_prob = t.nn.functional.log_softmax(inputs, dim=-1) |
| target_prob = t.nn.functional.softmax(targets, dim=-1) |
| |
| |
| kl_div = t.nn.functional.kl_div( |
| input_log_prob, |
| target_prob, |
| reduction=reduction |
| ) |
| |
| return kl_div |
|
|
| def str_to_torch_dtype(dtype_str: str) -> t.dtype: |
| dtype_mapping = { |
| "float16": t.float16, |
| "bfloat16": t.bfloat16, |
| "float32": t.float32, |
| "float64": t.float64, |
| "int8": t.int8, |
| "int16": t.int16, |
| "int32": t.int32, |
| "int64": t.int64, |
| "uint8": t.uint8, |
| "bool": t.bool, |
| } |
|
|
| if dtype_str not in dtype_mapping: |
| raise ValueError(f"Invalid dtype string: {dtype_str}. Supported dtypes are: {list(dtype_mapping.keys())}") |
|
|
| return dtype_mapping[dtype_str] |
|
|
| def split_data( |
| dataset: TensorDataset, |
| train_split: float = 0.8, |
| seed: int = 42, |
| ) -> Tuple[TensorDataset, TensorDataset]: |
| num_samples = len(dataset) |
| train_size = int(train_split * num_samples) |
| val_size = num_samples - train_size |
|
|
| t.manual_seed(seed) |
| train_dataset, val_dataset = random_split( |
| dataset, [train_size, val_size] |
| ) |
| |
| return train_dataset, val_dataset |