원 논문의 UniTTA를 이미지용 공식 구현 그대로 옮긴 것이 아님...!
- TabTransformer
- tabular input
- DESI inference
- streaming / online test-time adaptation
에 맞게 UniTTA의 핵심 아이디어를 가져와서 tabular feature-space 방식으로 재구성
모델 weight를 직접 업데이트하지 않고,
feature space에서
- domain-aware normalization
- class-balanced domain statistics
- temporal correlation 기반 feature adjustment
를 수행해서 target domain에 적응시키는 방식
코드 플롯
test stream을 따라가면서 target domain의 feature distribution을 추적하고,
그 정보를 이용해 feature representation을 보정한 뒤 최종 분류를 수행
- 학습된 모델을 불러온다.
- 입력 CSV를 읽는다.
- 배치 단위로 데이터를 순서대로 본다.
- 현재 배치가 어떤 “domain cluster”에 가까운지 판단한다.
- 그 domain에서 confident한 pseudo-label 샘플들의 class별 통계를 모은다.
- 그 통계로 feature를 재정렬한다.
- 이전 배치들과의 연속성도 반영해서 feature를 추가 보정한다.
- 그 보정된 feature로 최종 분류한다.
- 결과를 CSV로 저장한다.
모델 파라미터를 바꾸는 adaptation이 아니라,
feature-space adaptation
UniTTA 해석 구현법
UniTTA-style로 구현한 핵심 두 축
BDN-style
Balanced Domain Normalization 비슷한 역할을 하는 부분
각 domain 안에서 class별 feature 통계를 따로 저장한 뒤,
그 class별 통계를 평균내서 “class imbalance에 덜 치우친 domain 통계”를 만듦
즉 단순히 현재 배치 전체 평균만 쓰는 게 아니라,
domain 내부에서 confident한 pseudo-label을 이용해 class-wise mean/var를 유지하고,
그 평균을 balanced domain stat처럼 사용
COFA-style
Correlated Feature Adaptation 비슷한 역할을 하는 부분
이전 batch들의 평균 feature를 temporal memory에 저장해 두고,
현재 feature가 과거 feature와 얼마나 비슷한지에 따라
이전 context를 현재 feature에 섞음
즉 시간적으로 인접한 샘플들이 어느 정도 correlated하다는 가정 아래,
현재 sample representation을 직전 흐름과 더 일관되게 만드는 것
- 모델 가중치는 고정
- target feature distribution은 domain별로 추적
- domain/class 통계를 활용해 feature normalization
- 시간적 연속성을 활용해 feature correction
UniTTA-style 하이퍼파라미터
UNITTA_MAX_DOMAINS
최대 몇 개의 domain cluster를 유지할지
즉 target stream 전체를 하나의 분포로 보지 않고,
몇 개의 잠재 domain으로 나눠서 관리
UNITTA_DOMAIN_ASSIGN_TAU
현재 batch가 기존 domain 중 하나에 속하는지, 아니면 새 domain으로 볼지를 결정하는 민감도
값이 작으면 새 domain이 쉽게 생기고
값이 크면 대부분 기존 domain에 붙음
UNITTA_DOMAIN_MOMENTUM
domain centroid를 EMA처럼 부드럽게 업데이트할 때 쓰는 계수
UNITTA_STAT_MOMENTUM
domain-class 통계를 업데이트할 때 쓰는 EMA 계수
UNITTA_CONF_THRESHOLD
pseudo-label을 얼마나 믿을지 정하는 confidence 기준
즉 확률이 이 값 이상인 샘플만
“이 클래스에 속한다고 비교적 믿을 만하다”
라고 보고 domain-class 통계 업데이트에 사용
UNITTA_BLEND_SOURCE_STATS
source/global 통계와 target domain 통계를 얼마나 섞을지
이 값이 필요한 이유는
target domain 통계가 초반에는 불안정할 수 있기 때문
그래서 source stat을 일부 섞어서 안정화
UNITTA_USE_BALANCED_DOMAIN_NORM
BDN-style 정규화를 사용할지 여부
UNITTA_USE_COFA
COFA-style temporal feature adjustment를 사용할지 여부
UNITTA_COFA_ALPHA
과거 feature memory를 얼마나 강하게 반영할지 정함
UNITTA_TEMPORAL_WINDOW
최근 몇 개 batch의 평균 feature를 temporal memory로 유지할지 정함
UNITTA_CORR_POWER
현재 feature와 과거 feature가 얼마나 비슷한지 계산한 correlation gate를 몇 제곱해서 쓸지 정함
이 값이 크면 correlation이 높은 샘플에만 더 강하게 반응
UNITTA_LOGIT_TEMPERATURE
최종 logits를 softmax 전에 temperature scaling하는 역할
1.0이면 그대로 사용
UNITTA_DOMAIN_WARMUP_BATCHES
초기 몇 개 batch 동안은 domain adaptation을 바로 강하게 적용하지 않고 warmup하는 용도
UniTTAState 클래스
source_mean, source_var
초기 feature distribution을 source/global reference처럼 저장
domain_centroids
각 domain의 중심 벡터
현재 batch가 어떤 domain에 가까운지 판단할 때 사용
즉 target 데이터를 하나의 분포로 보지 않고,
여러 도메인 centroid를 두고 관리
domain_class_stats
domain_id -> class_id -> {mean, var, count}
즉 특정 domain 안에서,
class 0로 confident하게 예측된 샘플들의 feature 평균/분산
class 1로 confident하게 예측된 샘플들의 feature 평균/분산
class 2로 confident하게 예측된 샘플들의 feature 평균/분산
을 따로 저장
왜 이렇게 하냐면,
그냥 domain 전체 평균만 쓰면 class imbalance의 영향을 많이 받기 때문
예를 들어 domain 안에 Non 샘플이 훨씬 많으면,
domain 전체 평균이 사실상 Non 중심으로 쏠릴 수 있음
그래서 class별 통계를 따로 추적한 뒤 평균을 내서 balanced domain stats를 만듦
temporal_feat_memory / temporal_prob_memory
COFA-style temporal memory
최근 batch들의 평균 feature와 평균 확률을 저장
이전 문맥을 현재 batch feature adjustment에 반영하기 위한 메모리
UniTTAState 내부 메서드
initialize_source_stats
첫 batch의 feature mean/var를 source-like reference로 저장
이 구현에서는 학습 source 데이터 통계를 별도 파일로 저장하지 않았으므로,
초기 batch를 reference처럼 쓰는 방식
assign_domain
이 함수는 현재 batch가 어느 domain에 속하는지 결정
- 현재 batch feature mean을 구한다.
- 기존 domain centroids와 cosine similarity를 계산한다.
- 가장 가까운 domain을 찾는다.
- 거리가 너무 크면 새 domain을 만든다.
- 아니면 기존 domain centroid를 EMA로 업데이트한다.
즉 streaming target batch를 보면서,
현재 배치가 어떤 latent domain cluster에 속하는지 점진적으로 구성
이게 중요한 이유는,
모든 target sample을 하나의 통계로 처리하는 것보다,
여러 domain으로 나눠서 관리하는 게 더 유연할 수 있기 때문
update_domain_class_stats
현재 batch에서 confident pseudo-label을 가진 샘플들만 골라,
domain x class 통계를 업데이트
예를 들어 current domain = 2이고,
그 안에서 class 1 confidence 높은 샘플들이 있으면,
domain 2의 class 1 mean/var를 EMA 방식으로 갱신
즉 이 함수는 “domain 내부 클래스별 feature prototype/statistics bank”를 쌓는 역할
get_balanced_domain_stats
BDN-style의 핵심 구현
특정 domain에서 class별 mean/var를 모은 뒤,
그 class mean들끼리 다시 평균을 내서
balanced domain mean/var를 만듦
즉 어떤 클래스가 샘플 수가 훨씬 많아도,
그 클래스만 domain 통계를 지배하지 않도록 class별로 균형 있게 평균을 내는 것
그다음 source/global stat도 일부 섞음
이 blending은 초반 noisy한 target stat 때문에 feature normalization이 불안정해지는 걸 막기 위한 안정화 장치
cofa_adjust
COFA-style 구현
현재 feature와 최근 memory feature 평균의 cosine correlation을 구함
그다음 correlation이 높을수록
이전 feature context를 더 많이 섞음
즉 현재 feature가 이전 흐름과 연속성이 높다면,
그 이전 문맥을 활용해 representation을 부드럽게 보정
adjusted = (1 - gate) * current + gate * previous_context
즉 무조건 과거를 섞는 게 아니라,
현재와 과거가 비슷할수록 더 많이 반영
update_temporal_memory
현재 batch의 평균 feature와 평균 probability를 memory에 저장
이 memory는 COFA에서 다음 batch들을 보정하는 reference로 쓰임
run_inference_unitta
UniTTA-style 추론을 수행
모델 복사와 eval
model = copy.deepcopy(model).to(device)
model.eval()
원본 모델을 보존하기 위해 deepcopy를 하고,
weight update는 하지 않으므로 eval mode를 유지
즉 CoTTA처럼 optimizer나 backward는 없
state 초기화
처음 batch에서 feature dimension을 확인한 뒤 UniTTAState 생성
그리고 첫 batch feature를 기반으로 source/global stats를 초기화
raw feature / raw prediction
feats = model.forward_features(X_batch)
logits_raw = model.forward_logits_from_features(feats)
probs_raw = F.softmax(logits_raw, dim=1)
confs_raw, preds_raw = torch.max(probs_raw, dim=1)
즉 먼저 원래 모델 기준 feature와 예측을 얻음
raw prediction 두 가지 용도
- confident pseudo-label 추출
- domain-class stats 업데이트
domain assignment
batch_feat_mean = feats.mean(dim=0)
domain_id = state.assign_domain(batch_feat_mean)
현재 batch가 어떤 domain에 속하는지 정함
즉 target stream을 latent domain sequence처럼 관리
domain x class stats 업데이트
state.update_domain_class_stats(...)
현재 batch의 confident pseudo-labeled feature를 사용하여
현재 domain의 class별 통계를 갱신
즉 domain 내부에서 class-aware statistics bank를 점점 정교하게 쌓는 것
BDN-style feature normalization
dom_mean, dom_var, valid = state.get_balanced_domain_stats(domain_id)
현재 domain의 balanced mean/var를 가져옴
그 후 source stat 기준 whitening 후,
현재 domain stat 기준 recoloring을 진행
feats_adapt = (feats_adapt - src_mean.unsqueeze(0)) / torch.sqrt(src_var.unsqueeze(0) + EPS)
feats_adapt = feats_adapt * torch.sqrt(dom_var.unsqueeze(0) + EPS) + dom_mean.unsqueeze(0)
현재 feature를 source-like 통계 기준으로 정규화한 다음,
target domain의 balanced 통계에 맞게 다시 스케일/이동시키는 것
즉 domain-aware normalization을 feature space에서 수행하는 느낌?
COFA-style feature adjustment
feats_adapt = state.cofa_adjust(feats_adapt, probs_raw)
이제 방금 domain-normalized feature를
이전 temporal memory와의 correlation을 기준으로 한 번 더 보정
지금 batch가 최근 흐름과 비슷하다면, 최근 feature context를 조금 섞는 거
최종 prediction
logits_final = model.forward_logits_from_features(feats_adapt)
probs_final = F.softmax(logits_final, dim=1)
preds_final = torch.argmax(probs_final, dim=1)
즉 weight는 그대로 두고,
보정된 feature representation만 classifier head에 넣어 최종 예측
모델은 바꾸지 않고,
입력 representation을 target domain에 맞게 정렬한 뒤 분류
temporal memory update
state.update_temporal_memory(feats_adapt, probs_final)
현재 batch의 평균 feature / 평균 확률을 temporal memory에 저장
다음 batch에서 COFA가 이걸 참조
즉 streaming adaptation의 순환 구조
코드핵심
TabTransformer
tabular feature를 Transformer 기반 latent representation으로 바꾸는 분류 모델
UniTTAState
target stream adaptation에 필요한 모든 domain/stat/memory 상태 저장소
assign_domain
현재 batch가 어느 latent domain에 속하는지 결정
update_domain_class_stats
domain 내부 class별 feature 통계 누적
get_balanced_domain_stats
class imbalance를 줄인 balanced domain mean/var 생성
cofa_adjust
과거 temporal context를 현재 feature에 correlation 기반으로 반영
run_inference_unitta
모델 weight는 그대로 두고 feature-space에서 online target adaptation 수행
CoTTA <-> UniTTA-style 구현
CoTTA
- student weight 업데이트
- teacher EMA
- augmentation averaging
- stochastic restoration
- 즉 parameter adaptation
“모델을 바꾸는 방식”
UniTTA-style
- 모델 weight 고정
- feature/domain stats 추적
- balanced domain normalization
- temporal feature correction
- 즉 feature/statistics adaptation
“표현을 바꾸는 방식”
학습된 TabTransformer의 weight는 고정한 채,
target stream에서 domain별·class별 feature 통계와 temporal memory를 점진적으로 구축하고,
이를 이용해 feature representation을 domain-aware / correlation-aware하게 보정한 뒤
최종 분류를 수행하는 tabular continual TTA 구현
import os
import json
import copy
import random
from collections import defaultdict, deque
import pandas as pd
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
# ============================================================
# 0. Reproducibility
# ============================================================
GLOBAL_SEED = 42
random.seed(GLOBAL_SEED)
np.random.seed(GLOBAL_SEED)
torch.manual_seed(GLOBAL_SEED)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(GLOBAL_SEED)
# ============================================================
# 1. 경로 설정
# ============================================================
CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = os.path.abspath(os.path.join(CURRENT_DIR, "../../"))
MODEL_REL_PATH = "output/model/deep_learning/TabTransformer_multiseed/seed_12/TabTransformer_seed12.pkl"
BEST_PARAM_REL_PATH = "output/evaluation/deep_learning/TabTransformer_multiseed/seed_12/best_param/TabTransformer_best_param.json"
DATA_REL_PATH = "data/preprocessed/DESI_preprocessed.csv"
OUTPUT_REL_DIR = "output/inference"
TTA_MODE = "UNITTA"
OUTPUT_FILENAME = "DESI_TabTransformer_UNITTA_seed12.csv"
MODEL_PATH = os.path.join(PROJECT_ROOT, MODEL_REL_PATH)
BEST_PARAM_PATH = os.path.join(PROJECT_ROOT, BEST_PARAM_REL_PATH)
INPUT_CSV_PATH = os.path.join(PROJECT_ROOT, DATA_REL_PATH)
OUTPUT_DIR = os.path.join(PROJECT_ROOT, OUTPUT_REL_DIR)
OUTPUT_CSV_PATH = os.path.join(OUTPUT_DIR, OUTPUT_FILENAME)
# ============================================================
# 2. 기본 설정
# ============================================================
CLASS_NAMES = ["Non", "Pre", "Post"]
NUM_CLASSES = len(CLASS_NAMES)
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
EPS = 1e-8
FEATURES = [
"StellarMass", "AbsMag_g", "AbsMag_r", "AbsMag_i", "AbsMag_z",
"color_gr", "color_gi", "EffectiveRadius", "VelocityDispersion",
"Metallicity", "SFR", "BulgeMass"
]
AUX_COLS = ["P_NOMERGER", "RA", "DEC", "REDSHIFT"]
# ------------------------------------------------------------
# UniTTA-style 하이퍼파라미터
# ------------------------------------------------------------
UNITTA_BATCH_SIZE = 256
# domain bank
UNITTA_MAX_DOMAINS = 4
UNITTA_DOMAIN_ASSIGN_TAU = 0.35 # 새 domain 생성 민감도
UNITTA_DOMAIN_MOMENTUM = 0.95 # domain centroid EMA
UNITTA_STAT_MOMENTUM = 0.90 # class/domain stat EMA
# BDN-style
UNITTA_CONF_THRESHOLD = 0.70 # pseudo-label 신뢰도 기준
UNITTA_USE_BALANCED_DOMAIN_NORM = True
UNITTA_BLEND_SOURCE_STATS = 0.30 # source/global stats blending
# COFA-style
UNITTA_USE_COFA = True
UNITTA_COFA_ALPHA = 0.35 # 이전 feature 반영 강도 기본값
UNITTA_TEMPORAL_WINDOW = 8 # 최근 feature memory 길이
UNITTA_CORR_POWER = 1.0 # correlation gate exponent
# 안정화
UNITTA_LOGIT_TEMPERATURE = 1.0
UNITTA_DOMAIN_WARMUP_BATCHES = 1
# ============================================================
# 3. Dataset
# ============================================================
class GalaxyInferenceDataset(Dataset):
def __init__(self, X: np.ndarray):
self.X = torch.tensor(X, dtype=torch.float32)
def __len__(self):
return len(self.X)
def __getitem__(self, idx):
return self.X[idx]
# ============================================================
# 4. Model Definition
# ============================================================
class NumericalEmbedding(nn.Module):
def __init__(self, num_features: int, d_token: int):
super().__init__()
self.num_features = num_features
self.embeddings = nn.ModuleList(
[nn.Linear(1, d_token) for _ in range(num_features)]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
tokens = []
for i in range(self.num_features):
feature_slice = x[:, i:i + 1]
token = self.embeddings[i](feature_slice)
tokens.append(token)
return torch.stack(tokens, dim=1)
class TabTransformer(nn.Module):
def __init__(
self,
num_features: int,
d_token: int,
n_layers: int,
n_heads: int,
d_ffn_factor: float,
dropout: float,
num_classes: int = 3,
):
super().__init__()
self.embedding = NumericalEmbedding(num_features, d_token)
d_ffn = int(d_token * d_ffn_factor)
encoder_layer = nn.TransformerEncoderLayer(
d_model=d_token,
nhead=n_heads,
dim_feedforward=d_ffn,
dropout=dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
self.transformer = nn.TransformerEncoder(
encoder_layer,
num_layers=n_layers,
enable_nested_tensor=False,
)
self.flatten_dim = num_features * d_token
self.mlp_head = nn.Sequential(
nn.LayerNorm(self.flatten_dim),
nn.Dropout(dropout),
nn.Linear(self.flatten_dim, d_token),
nn.GELU(),
nn.Linear(d_token, num_classes),
)
def forward_features(self, x: torch.Tensor) -> torch.Tensor:
x = self.embedding(x)
x = self.transformer(x)
x = x.reshape(x.size(0), -1)
return x
def forward_logits_from_features(self, feats: torch.Tensor) -> torch.Tensor:
return self.mlp_head(feats)
def forward(self, x: torch.Tensor) -> torch.Tensor:
feats = self.forward_features(x)
logits = self.forward_logits_from_features(feats)
return logits
# ============================================================
# 5. Utils
# ============================================================
def load_model(model_path: str, best_param_path: str, num_features: int, device: str):
if not os.path.exists(best_param_path):
raise FileNotFoundError(f"Best param file not found: {best_param_path}")
if not os.path.exists(model_path):
raise FileNotFoundError(f"Model checkpoint not found: {model_path}")
with open(best_param_path, "r", encoding="utf-8") as f:
best_params = json.load(f)
model = TabTransformer(
num_features=num_features,
d_token=best_params["d_token"],
n_layers=best_params["n_layers"],
n_heads=best_params["n_heads"],
d_ffn_factor=best_params["d_ffn_factor"],
dropout=best_params["dropout"],
num_classes=NUM_CLASSES,
).to(device)
state_dict = torch.load(model_path, map_location=device)
model.load_state_dict(state_dict)
model.eval()
return model, best_params
def l2_normalize(x: torch.Tensor, dim: int = 1, eps: float = 1e-8) -> torch.Tensor:
return x / (torch.norm(x, dim=dim, keepdim=True) + eps)
def compute_batch_mean_var(x: torch.Tensor):
mean = x.mean(dim=0)
var = x.var(dim=0, unbiased=False)
return mean, var
def ema_update(old: torch.Tensor, new: torch.Tensor, momentum: float) -> torch.Tensor:
return momentum * old + (1.0 - momentum) * new
# ============================================================
# 6. UniTTA-style State
# ============================================================
class UniTTAState:
"""
Tabular feature space에서 UniTTA-style state를 관리
구성:
- source/global stats
- domain centroid bank
- domain x class feature stats
- COFA temporal memory
"""
def __init__(
self,
feat_dim: int,
num_classes: int,
max_domains: int = 4,
domain_assign_tau: float = 0.35,
domain_momentum: float = 0.95,
stat_momentum: float = 0.90,
conf_threshold: float = 0.70,
blend_source_stats: float = 0.30,
temporal_window: int = 8,
cofa_alpha: float = 0.35,
corr_power: float = 1.0,
):
self.feat_dim = feat_dim
self.num_classes = num_classes
self.max_domains = max_domains
self.domain_assign_tau = domain_assign_tau
self.domain_momentum = domain_momentum
self.stat_momentum = stat_momentum
self.conf_threshold = conf_threshold
self.blend_source_stats = blend_source_stats
self.temporal_window = temporal_window
self.cofa_alpha = cofa_alpha
self.corr_power = corr_power
self.source_mean = None
self.source_var = None
self.domain_centroids = [] # list[(D,)]
self.domain_counts = [] # list[int]
self.num_seen_batches = 0
# domain -> class -> dict(mean, var, count)
self.domain_class_stats = defaultdict(dict)
# COFA memory
self.temporal_feat_memory = deque(maxlen=temporal_window)
self.temporal_prob_memory = deque(maxlen=temporal_window)
def initialize_source_stats(self, feats: torch.Tensor):
mean, var = compute_batch_mean_var(feats)
self.source_mean = mean.detach().clone()
self.source_var = var.detach().clone()
def assign_domain(self, batch_feat_mean: torch.Tensor) -> int:
"""
현재 batch를 가장 가까운 domain에 할당하거나 새 domain 생성
cosine distance 기반
"""
batch_feat_mean = batch_feat_mean.detach()
if len(self.domain_centroids) == 0:
self.domain_centroids.append(batch_feat_mean.clone())
self.domain_counts.append(1)
return 0
centroids = torch.stack(self.domain_centroids, dim=0) # (K, D)
x = l2_normalize(batch_feat_mean.unsqueeze(0), dim=1) # (1, D)
c = l2_normalize(centroids, dim=1) # (K, D)
sims = (x @ c.t()).squeeze(0) # (K,)
best_idx = int(torch.argmax(sims).item())
best_sim = float(sims[best_idx].item())
best_dist = 1.0 - best_sim
if (best_dist > self.domain_assign_tau) and (len(self.domain_centroids) < self.max_domains):
self.domain_centroids.append(batch_feat_mean.clone())
self.domain_counts.append(1)
return len(self.domain_centroids) - 1
self.domain_centroids[best_idx] = ema_update(
self.domain_centroids[best_idx],
batch_feat_mean,
self.domain_momentum,
)
self.domain_counts[best_idx] += 1
return best_idx
def update_domain_class_stats(
self,
domain_id: int,
feats: torch.Tensor,
preds: torch.Tensor,
confs: torch.Tensor,
):
"""
pseudo-label confident sample들만 사용해 domain x class 통계 업데이트
"""
for c in range(self.num_classes):
mask = (preds == c) & (confs >= self.conf_threshold)
if mask.sum().item() == 0:
continue
feats_c = feats[mask]
mean_c, var_c = compute_batch_mean_var(feats_c)
if c not in self.domain_class_stats[domain_id]:
self.domain_class_stats[domain_id][c] = {
"mean": mean_c.detach().clone(),
"var": var_c.detach().clone(),
"count": int(mask.sum().item()),
}
else:
prev = self.domain_class_stats[domain_id][c]
prev["mean"] = ema_update(prev["mean"], mean_c.detach(), self.stat_momentum)
prev["var"] = ema_update(prev["var"], var_c.detach(), self.stat_momentum)
prev["count"] += int(mask.sum().item())
def get_balanced_domain_stats(self, domain_id: int):
"""
BDN-style:
domain 내부에서 class별 mean/var를 구한 뒤 class 평균으로 balanced domain stats 생성
"""
if domain_id not in self.domain_class_stats:
return self.source_mean, self.source_var, False
cls_stats = self.domain_class_stats[domain_id]
if len(cls_stats) == 0:
return self.source_mean, self.source_var, False
means = []
vars_ = []
for c in sorted(cls_stats.keys()):
means.append(cls_stats[c]["mean"])
vars_.append(cls_stats[c]["var"])
dom_mean = torch.stack(means, dim=0).mean(dim=0)
dom_var = torch.stack(vars_, dim=0).mean(dim=0)
# source/global stats와 blending하여 안정화
if self.source_mean is not None and self.source_var is not None:
dom_mean = (1.0 - self.blend_source_stats) * dom_mean + self.blend_source_stats * self.source_mean
dom_var = (1.0 - self.blend_source_stats) * dom_var + self.blend_source_stats * self.source_var
return dom_mean, dom_var, True
def cofa_adjust(self, feats: torch.Tensor, probs: torch.Tensor):
"""
COFA-style:
최근 temporal memory의 평균 feature를 참조해 현재 feature를 보정
correlation이 높을수록 이전 context를 더 반영
"""
if len(self.temporal_feat_memory) == 0:
return feats
prev_feat = torch.stack(list(self.temporal_feat_memory), dim=0).mean(dim=0) # (D,)
prev_feat = prev_feat.unsqueeze(0).expand_as(feats) # (B, D)
feats_n = l2_normalize(feats, dim=1, eps=EPS)
prev_n = l2_normalize(prev_feat, dim=1, eps=EPS)
corr = (feats_n * prev_n).sum(dim=1, keepdim=True).clamp(min=0.0, max=1.0)
gate = self.cofa_alpha * (corr ** self.corr_power)
adjusted = (1.0 - gate) * feats + gate * prev_feat
return adjusted
def update_temporal_memory(self, feats: torch.Tensor, probs: torch.Tensor):
"""
batch 평균 representation을 저장
"""
feat_mean = feats.mean(dim=0).detach()
prob_mean = probs.mean(dim=0).detach()
self.temporal_feat_memory.append(feat_mean)
self.temporal_prob_memory.append(prob_mean)
# ============================================================
# 7. UniTTA-style Inference
# ============================================================
def run_inference_unitta(
model: nn.Module,
loader: DataLoader,
device: str,
max_domains: int = 4,
domain_assign_tau: float = 0.35,
domain_momentum: float = 0.95,
stat_momentum: float = 0.90,
conf_threshold: float = 0.70,
blend_source_stats: float = 0.30,
use_balanced_domain_norm: bool = True,
use_cofa: bool = True,
cofa_alpha: float = 0.35,
temporal_window: int = 8,
corr_power: float = 1.0,
logit_temperature: float = 1.0,
warmup_batches: int = 1,
):
"""
TabTransformer용 UniTTA-style 구현
아이디어:
1) feature 추출
2) batch mean 기반 domain assignment
3) confident pseudo-label 기반 domain x class stats 업데이트
4) class-balanced domain stats로 feature normalization (BDN-style)
5) 이전 temporal feature memory를 활용한 feature correction (COFA-style)
6) classifier head로 최종 예측
주의:
- 모델 weight는 직접 업데이트하지 않음
- feature-space adaptation 중심
"""
model = copy.deepcopy(model).to(device)
model.eval()
state = None
all_probs = []
all_preds = []
with torch.no_grad():
for batch_idx, X_batch in enumerate(loader):
X_batch = X_batch.to(device)
# ----------------------------------------
# 1) feature / original prediction
# ----------------------------------------
feats = model.forward_features(X_batch) # (B, D)
if state is None:
state = UniTTAState(
feat_dim=feats.shape[1],
num_classes=NUM_CLASSES,
max_domains=max_domains,
domain_assign_tau=domain_assign_tau,
domain_momentum=domain_momentum,
stat_momentum=stat_momentum,
conf_threshold=conf_threshold,
blend_source_stats=blend_source_stats,
temporal_window=temporal_window,
cofa_alpha=cofa_alpha,
corr_power=corr_power,
)
state.initialize_source_stats(feats)
logits_raw = model.forward_logits_from_features(feats) / logit_temperature
probs_raw = F.softmax(logits_raw, dim=1)
confs_raw, preds_raw = torch.max(probs_raw, dim=1)
# ----------------------------------------
# 2) domain assignment
# ----------------------------------------
batch_feat_mean = feats.mean(dim=0)
domain_id = state.assign_domain(batch_feat_mean)
# ----------------------------------------
# 3) domain x class stats 업데이트
# ----------------------------------------
state.update_domain_class_stats(
domain_id=domain_id,
feats=feats,
preds=preds_raw,
confs=confs_raw,
)
# ----------------------------------------
# 4) BDN-style feature normalization
# ----------------------------------------
feats_adapt = feats
if use_balanced_domain_norm and batch_idx >= warmup_batches:
dom_mean, dom_var, valid = state.get_balanced_domain_stats(domain_id)
if valid:
# source/global stats 기준 feature whitening -> recoloring
src_mean = state.source_mean
src_var = state.source_var
feats_adapt = (feats_adapt - src_mean.unsqueeze(0)) / torch.sqrt(src_var.unsqueeze(0) + EPS)
feats_adapt = feats_adapt * torch.sqrt(dom_var.unsqueeze(0) + EPS) + dom_mean.unsqueeze(0)
# ----------------------------------------
# 5) COFA-style correlated feature adaptation
# ----------------------------------------
if use_cofa and batch_idx >= warmup_batches:
feats_adapt = state.cofa_adjust(feats_adapt, probs_raw)
# ----------------------------------------
# 6) final prediction
# ----------------------------------------
logits_final = model.forward_logits_from_features(feats_adapt) / logit_temperature
probs_final = F.softmax(logits_final, dim=1)
preds_final = torch.argmax(probs_final, dim=1)
all_probs.append(probs_final.cpu().numpy())
all_preds.append(preds_final.cpu().numpy())
# ----------------------------------------
# 7) temporal memory update
# ----------------------------------------
state.update_temporal_memory(feats_adapt, probs_final)
all_probs = np.concatenate(all_probs, axis=0)
all_preds = np.concatenate(all_preds, axis=0)
return all_preds, all_probs
# ============================================================
# 8. Main
# ============================================================
def main():
print("=== UniTTA-style Inference Started ===")
print(f"Project Root: {PROJECT_ROOT}")
print(f"Using device: {DEVICE}")
print(f"TTA mode: {TTA_MODE}")
# --------------------------------------------------------
# 1) 데이터 로드
# --------------------------------------------------------
if not os.path.exists(INPUT_CSV_PATH):
raise FileNotFoundError(f"Input data not found at: {INPUT_CSV_PATH}")
print(f"Loading data from: {DATA_REL_PATH}")
df = pd.read_csv(INPUT_CSV_PATH)
print(f"Data shape: {df.shape}")
missing_feature_cols = [col for col in FEATURES if col not in df.columns]
if missing_feature_cols:
raise ValueError(f"Missing features in input data: {missing_feature_cols}")
missing_aux_cols = [col for col in AUX_COLS if col not in df.columns]
if missing_aux_cols:
raise ValueError(f"Missing aux columns in input data: {missing_aux_cols}")
X = df[FEATURES].values.astype(np.float32)
if np.isnan(X).any():
raise ValueError("Input features contain NaN values.")
if np.isinf(X).any():
raise ValueError("Input features contain inf values.")
num_features = X.shape[1]
print(f"Extracted {num_features} features for inference.")
# --------------------------------------------------------
# 2) 모델 로드
# --------------------------------------------------------
print(f"Loading model from: {MODEL_REL_PATH}")
model, best_params = load_model(
model_path=MODEL_PATH,
best_param_path=BEST_PARAM_PATH,
num_features=num_features,
device=DEVICE,
)
print("Model loaded successfully.")
print("Best params:", best_params)
# --------------------------------------------------------
# 3) DataLoader
# --------------------------------------------------------
infer_ds = GalaxyInferenceDataset(X)
infer_loader = DataLoader(
infer_ds,
batch_size=UNITTA_BATCH_SIZE,
shuffle=False,
)
# --------------------------------------------------------
# 4) UniTTA-style 실행
# --------------------------------------------------------
predictions, probabilities = run_inference_unitta(
model=model,
loader=infer_loader,
device=DEVICE,
max_domains=UNITTA_MAX_DOMAINS,
domain_assign_tau=UNITTA_DOMAIN_ASSIGN_TAU,
domain_momentum=UNITTA_DOMAIN_MOMENTUM,
stat_momentum=UNITTA_STAT_MOMENTUM,
conf_threshold=UNITTA_CONF_THRESHOLD,
blend_source_stats=UNITTA_BLEND_SOURCE_STATS,
use_balanced_domain_norm=UNITTA_USE_BALANCED_DOMAIN_NORM,
use_cofa=UNITTA_USE_COFA,
cofa_alpha=UNITTA_COFA_ALPHA,
temporal_window=UNITTA_TEMPORAL_WINDOW,
corr_power=UNITTA_CORR_POWER,
logit_temperature=UNITTA_LOGIT_TEMPERATURE,
warmup_batches=UNITTA_DOMAIN_WARMUP_BATCHES,
)
confidences = np.max(probabilities, axis=1)
# --------------------------------------------------------
# 5) 결과 정리
# --------------------------------------------------------
result_df = df[AUX_COLS].copy()
ordered_aux = ["RA", "DEC", "REDSHIFT", "P_NOMERGER"]
result_df = result_df[ordered_aux]
result_df["prediction"] = predictions
result_df["prediction_label"] = [CLASS_NAMES[p] for p in predictions]
result_df["confidence"] = confidences
result_df["non_confidence"] = probabilities[:, 0]
result_df["pre_confidence"] = probabilities[:, 1]
result_df["post_confidence"] = probabilities[:, 2]
# --------------------------------------------------------
# 6) 저장
# --------------------------------------------------------
os.makedirs(OUTPUT_DIR, exist_ok=True)
result_df.to_csv(OUTPUT_CSV_PATH, index=False)
print("=== UniTTA-style Inference Completed Successfully ===")
print(f"Saved file path: {OUTPUT_CSV_PATH}")
print(result_df.head())
if __name__ == "__main__":
main()'Club > 졸업 연구 | 멀티모달 AI를 이용한 은하 병합 단계 분류' 카테고리의 다른 글
| 🔳 TACT 적용 (TabTransformer) (0) | 2026.04.17 |
|---|---|
| 🔳 LATTA 적용 (TabTransformer) (0) | 2026.04.17 |
| 🔳 CoTTA 적용 (TabTransformer) (0) | 2026.04.15 |
| TabTransformer inference + TTA 분석 (0) | 2026.04.14 |
| GradientBoosting 코드 검증 (0) | 2026.04.06 |