import numpy as np
from numpy.linalg import norm
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler

DATA_PATH = "path to train features as .npz"
SAVE_PATH = "save the pseudolabeled data as .npz"

CLS_DIM = 1536            
TARGET_PRECISION = 0.99   
CALIB_FRACTION = 0.40     
MIN_CALIB_PER_CLASS = 8   
N_ROUNDS = 3
SEED = 42

if CLS_DIM is None:
    raise ValueError("Set CLS_DIM to match your other scripts.")

np.random.seed(SEED)

data = np.load(DATA_PATH, allow_pickle=True)
X_full_raw = data["features"].astype(np.float32)
y_full = data["labels"].astype(int)
image_paths = data["paths"]


def split_normalize(X, cls_dim):
    X_cls = X[:, :cls_dim]
    X_mean = X[:, cls_dim:]
    X_cls_n = X_cls / (norm(X_cls, axis=1, keepdims=True) + 1e-12)
    X_mean_n = X_mean / (norm(X_mean, axis=1, keepdims=True) + 1e-12)
    return np.concatenate([X_cls_n, X_mean_n], axis=1)


X_full_norm = split_normalize(X_full_raw, CLS_DIM)

N = len(X_full_norm)
indices = np.random.permutation(N)
n_labeled = int(0.3 * N)

labeled_idx = indices[:n_labeled]
unlabeled_idx = indices[n_labeled:]

X_l_raw = X_full_raw[labeled_idx]
X_u_raw = X_full_raw[unlabeled_idx]
X_l = X_full_norm[labeled_idx]
X_u = X_full_norm[unlabeled_idx]
print("Number of labeled samples   :", len(X_l))
print("Number of unlabeled samples :", len(X_u))
y_l = y_full[labeled_idx]
y_u_true = y_full[unlabeled_idx]  

paths_l = image_paths[labeled_idx]
paths_u = image_paths[unlabeled_idx]

classes = np.sort(np.unique(y_l))  
K = len(classes)

print("Gold labeled per class:", {int(c): int(np.sum(y_l == c)) for c in classes})


gold_train_idx, gold_calib_idx = train_test_split(
    np.arange(len(X_l)), test_size=CALIB_FRACTION, stratify=y_l, random_state=SEED
)

X_calib, y_calib = X_l[gold_calib_idx], y_l[gold_calib_idx]

print("Calibration set size per class:",
      {int(c): int(np.sum(y_calib == c)) for c in classes})

def fit_ordinal_classifier(X_train, y_train, scaler):
    
    X_train_s = scaler.transform(X_train)
    binary_clfs = []
    for k in range(K - 1):
        y_bin = (y_train > classes[k]).astype(int)
        clf = LogisticRegression(
            penalty="l2", C=0.5, class_weight="balanced",
            max_iter=2000, solver="lbfgs"
        )
        clf.fit(X_train_s, y_bin)
        binary_clfs.append(clf)
    return binary_clfs


def ordinal_class_probs(X, scaler, binary_clfs):
   
    X_s = scaler.transform(X)
    # P_gt[:, k] = P(y > classes[k])
    P_gt = np.stack([clf.predict_proba(X_s)[:, 1] for clf in binary_clfs], axis=1)

    P_class = np.zeros((len(X), K))
    P_class[:, 0] = 1 - P_gt[:, 0]
    for k in range(1, K - 1):
        P_class[:, k] = P_gt[:, k - 1] - P_gt[:, k]
    P_class[:, K - 1] = P_gt[:, K - 2]

    P_class = np.clip(P_class, 0, None)
    row_sums = P_class.sum(axis=1, keepdims=True)
    row_sums[row_sums == 0] = 1.0
    P_class = P_class / row_sums
    return P_class


def calibrate_thresholds(P_calib, y_calib_true, predicted_calib):

    thresholds = {}
    for ci, c in enumerate(classes):
        mask = predicted_calib == c
        n_c = mask.sum()

        if n_c < MIN_CALIB_PER_CLASS:
            thresholds[c] = None
            print(f"  Class {c}: only {n_c} calibration samples predicted "
                  f"this class (<{MIN_CALIB_PER_CLASS}) -- skipping, cannot calibrate reliably.")
            continue

        conf_c = P_calib[mask, ci]
        correct_c = (y_calib_true[mask] == c)

        order = np.argsort(-conf_c)
        conf_sorted = conf_c[order]
        correct_sorted = correct_c[order]

        cum_correct = np.cumsum(correct_sorted)
        cum_total = np.arange(1, len(correct_sorted) + 1)
        cum_precision = cum_correct / cum_total

        valid = np.where(cum_precision >= TARGET_PRECISION)[0]
        if len(valid) == 0:
            thresholds[c] = None
            best_prec = cum_precision.max() if len(cum_precision) > 0 else 0.0
            print(f"  Class {c}: target precision {TARGET_PRECISION} not reachable "
                  f"(best achievable={best_prec:.3f}) -- skipping this round.")
        else:
            best_cutoff = valid.max()
            tau = conf_sorted[best_cutoff]
            thresholds[c] = tau
            print(f"  Class {c}: threshold={tau:.4f}  "
                  f"(covers {best_cutoff + 1}/{n_c} calib samples at "
                  f"precision={cum_precision[best_cutoff]:.3f})")
    return thresholds


cur_X_train_idx_gold = gold_train_idx.copy()  
accum_pseudo_X_norm = []
accum_pseudo_y = []

final_pseudo_X_raw, final_pseudo_y, final_pseudo_paths, final_pseudo_conf = [], [], [], []
remaining_mask = np.ones(len(X_u), dtype=bool)

for round_i in range(1, N_ROUNDS + 1):
    print(f"\nRound {round_i} ")

    if accum_pseudo_X_norm:
        X_train = np.vstack([X_l[cur_X_train_idx_gold]] + accum_pseudo_X_norm)
        y_train = np.concatenate([y_l[cur_X_train_idx_gold]] + accum_pseudo_y)
    else:
        X_train = X_l[cur_X_train_idx_gold]
        y_train = y_l[cur_X_train_idx_gold]

    scaler = StandardScaler().fit(X_train)
    binary_clfs = fit_ordinal_classifier(X_train, y_train, scaler)

    P_calib = ordinal_class_probs(X_calib, scaler, binary_clfs)
    predicted_calib = classes[np.argmax(P_calib, axis=1)]
    print("Calibrating thresholds on held-out gold data:")
    thresholds = calibrate_thresholds(P_calib, y_calib, predicted_calib)

    pool_idx = np.where(remaining_mask)[0]
    if len(pool_idx) == 0:
        print("No unlabeled samples left. Stopping.")
        break

    X_pool = X_u[pool_idx]
    y_pool_true = y_u_true[pool_idx]  # diagnostics only

    P_pool = ordinal_class_probs(X_pool, scaler, binary_clfs)
    predicted_pool = classes[np.argmax(P_pool, axis=1)]
    conf_pool = P_pool.max(axis=1)

    accepted_mask = np.zeros(len(pool_idx), dtype=bool)
    for ci, c in enumerate(classes):
        tau = thresholds[c]
        if tau is None:
            continue
        sel = (predicted_pool == c) & (conf_pool >= tau)
        accepted_mask |= sel

    n_accepted = accepted_mask.sum()
    print(f"Accepted this round: {n_accepted}")
    if n_accepted == 0:
        print("No new confident samples at target precision. Stopping.")
        break

    for c in classes:
        cmask = accepted_mask & (predicted_pool == c)
        if cmask.sum() == 0:
            continue
        acc = np.mean(predicted_pool[cmask] == y_pool_true[cmask])
        print(f"  [diagnostic] Class {c}: selected={cmask.sum():4d}  "
              f"true precision on unlabeled={acc:.4f}")

    accepted_global_idx = pool_idx[accepted_mask]

    new_X_norm = X_u[accepted_global_idx]
    new_X_raw = X_u_raw[accepted_global_idx]
    new_y = predicted_pool[accepted_mask]
    new_paths = paths_u[accepted_global_idx]
    new_conf = conf_pool[accepted_mask]

    accum_pseudo_X_norm.append(new_X_norm)
    accum_pseudo_y.append(new_y)

    final_pseudo_X_raw.append(new_X_raw)
    final_pseudo_y.append(new_y)
    final_pseudo_paths.append(new_paths)
    final_pseudo_conf.append(new_conf)

    remaining_mask[accepted_global_idx] = False


if final_pseudo_X_raw:
    pseudo_X_raw = np.vstack(final_pseudo_X_raw)
    pseudo_y = np.concatenate(final_pseudo_y)
    pseudo_paths = np.concatenate(final_pseudo_paths)
    pseudo_conf = np.concatenate(final_pseudo_conf)
else:
    pseudo_X_raw = np.empty((0, X_l_raw.shape[1]), dtype=np.float32)
    pseudo_y = np.empty((0,), dtype=int)
    pseudo_paths = np.empty((0,), dtype=image_paths.dtype)
    pseudo_conf = np.empty((0,), dtype=np.float32)

num_selected = len(pseudo_y)
if num_selected > 0:
    path_to_true = {p: t for p, t in zip(paths_u, y_u_true)}
    true_for_pseudo = np.array([path_to_true[p] for p in pseudo_paths])
    precision = np.mean(pseudo_y == true_for_pseudo)
else:
    precision = 0.0

print("\nFINAL ")
print("Total pseudolabeled samples across all rounds:", num_selected)
print("Overall TRUE precision on unlabeled pool (diagnostic only):", round(precision, 4))
print("Class distribution of accepted pseudo-labels:",
      {int(c): int(np.sum(pseudo_y == c)) for c in classes})

X_combined_raw = np.vstack([X_l_raw, pseudo_X_raw])
y_combined = np.concatenate([y_l, pseudo_y])
combined_paths = np.concatenate([paths_l, pseudo_paths])

print("Final training size:", len(X_combined_raw))
print("Class distribution (combined):",
      {int(c): int(np.sum(y_combined == c)) for c in classes})

np.savez(
    SAVE_PATH,
    original_labeled_features=X_l_raw,
    original_labeled_labels=y_l,
    original_labeled_paths=paths_l,

    pseudolabeled_features=pseudo_X_raw,
    pseudolabeled_labels=pseudo_y,
    pseudolabeled_paths=pseudo_paths,
    pseudolabeled_confidence=pseudo_conf,

    combined_features=X_combined_raw,
    combined_labels=y_combined,
    combined_paths=combined_paths,

    pseudolabel_precision=precision,
    target_precision=TARGET_PRECISION,
)

print(f"\nSaved to {SAVE_PATH}")