본문 바로가기
  • 컴공생의 공부 일기
  • 공부보단 일기에 가까운 것 같은
  • 블로그
Club/졸업 연구 | 멀티모달 AI를 이용한 은하 병합 단계 분류

🔳 UniTTA 적용 (TabTransformer)

by 정람지 2026. 4. 17.

원 논문의 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을 보정한 뒤 최종 분류를 수행

 

  1. 학습된 모델을 불러온다.
  2. 입력 CSV를 읽는다.
  3. 배치 단위로 데이터를 순서대로 본다.
  4. 현재 배치가 어떤 “domain cluster”에 가까운지 판단한다.
  5. 그 domain에서 confident한 pseudo-label 샘플들의 class별 통계를 모은다.
  6. 그 통계로 feature를 재정렬한다.
  7. 이전 배치들과의 연속성도 반영해서 feature를 추가 보정한다.
  8. 그 보정된 feature로 최종 분류한다.
  9. 결과를 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에 속하는지 결정

  1. 현재 batch feature mean을 구한다.
  2. 기존 domain centroids와 cosine similarity를 계산한다.
  3. 가장 가까운 domain을 찾는다.
  4. 거리가 너무 크면 새 domain을 만든다.
  5. 아니면 기존 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 두 가지 용도

  1. confident pseudo-label 추출
  2. 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()