File size: 2,045 Bytes
a2ffd07
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
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'):
    # Convert logits to probabilities using softmax
    input_log_prob = t.nn.functional.log_softmax(inputs, dim=-1)
    target_prob = t.nn.functional.softmax(targets, dim=-1)
    
    # Calculate KL divergence
    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  # Remaining for test

    t.manual_seed(seed)
    train_dataset, val_dataset = random_split(
        dataset, [train_size, val_size]
    )
    
    return train_dataset, val_dataset # type: ignore