File size: 2,271 Bytes
64fd08f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
import torch.nn as nn
import pandas as pd


def get_optimizer(model, cnn_wd, transformer_wd, cnn_lr, transformer_lr, general_lr):
    """
    Helper function to set different weight decay for CNN and transformers.
    """
    
    # Figure out which parameters need decay
    decay, no_decay = set(), set()
    for mn, m in model.named_modules():
        for pn, p in m.named_parameters(recurse=False):
            fqn = f"{mn}.{pn}" if mn else pn
            if isinstance(m, (nn.LayerNorm, nn.BatchNorm1d, nn.BatchNorm2d)) or pn.endswith("bias"):
                no_decay.add(fqn)
            else:
                decay.add(fqn)
                
    # Set different weight decay for CNN and transformers
    param_dict = {pn: p for pn, p in model.named_parameters()}
    optimizer = torch.optim.AdamW(
        [
            {"params": [param_dict[n] for n in sorted(decay) if "transformer" not in n], "weight_decay": cnn_wd, "lr": cnn_lr},
            {"params": [param_dict[n] for n in sorted(decay) if "transformer" in n], "weight_decay": transformer_wd, "lr": transformer_lr},
            {"params": [param_dict[n] for n in sorted(no_decay)], "weight_decay": 0.0, "lr": general_lr},
        ]
    )
    
    return optimizer


def compute_class_weights(df: pd.DataFrame, label: str, dataset: str, all_scans: bool = False):
    """
    Compute the class weights for the given label.
    """
    
    # Make a copy to avoid SettingWithCopyWarning
    df = df.copy()
    
    # Map values to range [0, 2]
    if dataset == "odelia":
        odelia_mapper = {"normal": 0, "benign": 1, "malignant": 2}
        df["breast_label"] = df["breast_label"].map(odelia_mapper)

    # Only select a single scan per patient
    if all_scans:
        df = df[df["scan_number"] == 1]

    # Get counts for how often each class occurs
    class_counts = df[label].value_counts().sort_index()
    
    # Account for multiclass and binary classification
    if dataset == "odelia":
        class_weights = class_counts.max() / class_counts
        return torch.tensor(class_weights.values, dtype=torch.float32, device="cuda")
    else:
        class_weights = class_counts[0] / class_counts[1]
        return torch.tensor(class_weights, dtype=torch.float32, device="cuda")