# GOOGLE COLAB- MOUNT DRIVE


def maybe_mount_drive():
    try:
        import google.colab
        from google.colab import drive
        drive.mount("/content/drive")
        print("[Drive] Mounted at /content/drive")
        return True
    except Exception:
        print("[Drive] Not in Colab -- using local paths")
        return False

IN_COLAB = maybe_mount_drive()



# 1..... IMPORTS

import os, time, json, csv, math, random
from io import StringIO
from datetime import datetime
from zoneinfo import ZoneInfo

import numpy as np
import tensorflow as tf
from sklearn.cluster import MiniBatchKMeans
from sklearn.metrics import (confusion_matrix, classification_report,
                             f1_score, precision_recall_fscore_support)

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt





# 1.1) TIMEZONE HELPERS  (CST)

CENTRAL_TZ = ZoneInfo("America/Chicago")

def now_central():
    return datetime.now(CENTRAL_TZ)

def make_run_id():
    return "run_" + now_central().strftime("%Y_%m_%d_%H_%M_%S")





# 1.2) CIFAR-10 CLASS NAMES

CIFAR10_CLASSES = [
    "airplane", "automobile", "bird", "cat", "deer",
    "dog", "frog", "horse", "ship", "truck"
]




# 2) GPU VERIFICATION
# =============================================================================
def verify_gpu():
    print("\n" + "=" * 80)
    print("GPU VERIFICATION")
    print("=" * 80)
    gpus = tf.config.list_physical_devices("GPU")
    if not gpus:
        print("NO GPU DETECTED -- TensorFlow will run on CPU")
        print("   Colab fix: Runtime -> Change runtime type -> GPU")
    else:
        print(f"GPU DETECTED: {len(gpus)} GPU(s)")
        for i, gpu in enumerate(gpus):
            print(f"   GPU {i}: {gpu}")
        try:
            for gpu in gpus:
                tf.config.experimental.set_memory_growth(gpu, True)
            print("GPU memory growth enabled")
        except Exception as exc:
            print(f"[WARN] Could not set memory growth: {exc}")
    print(f"TensorFlow version : {tf.__version__}")
    print(f"Built with CUDA    : {tf.test.is_built_with_cuda()}")
    if gpus:
        try:
            with tf.device("/GPU:0"):
                _ = tf.matmul(
                    tf.constant([[1.0, 2.0], [3.0, 4.0]]),
                    tf.constant([[1.0, 1.0], [0.0, 1.0]]))
            print("GPU smoke-test passed")
        except Exception as exc:
            print(f"GPU smoke-test failed: {exc}")
    print("=" * 80 + "\n")

verify_gpu()




# ========================================================================
# 3 CONFIGURATION
# ======================================================================
CONFIG = {
    "dataset":            "CIFAR-10",
    "num_classes":        10,
    "image_shape":        (32, 32, 3),

    "k_per_class":        10,          
    "full_train_epochs":  1000,
    "distill_train_epochs": 1000,

    "base_alpha":         1e-4,

    "batch_size":         64,
    "learning_rate":      1e-3,
    "eta_min":            1e-5,
    "lr_gamma":           10.0,
    "probe_batch_size":   32,
    "probe_gradient_every_n_epochs": 10,

    "gradient_batch_size":      512,
    "gradient_sample_batches":  20,

    "kmeans_n_init":      3,
    "kmeans_max_iter":    100,
    "kmeans_batch_size":  2048,

    "random_seed":        42,

    "save_plots":             True,
    "save_arrays":            True,
    "save_distilled_grids":   True,
    "output_root":            None,
    "run_dir":                None,
    "_run_id":                None,
}

assert CONFIG["distill_train_epochs"] == CONFIG["full_train_epochs"], (
    f"distill_train_epochs ({CONFIG['distill_train_epochs']}) must equal "
    f"full_train_epochs ({CONFIG['full_train_epochs']})")

print("CONFIG  (edit only in section 3)")
print(f"   K per class           : {CONFIG['k_per_class']}")
print(f"   Full train epochs     : {CONFIG['full_train_epochs']}")
print(f"   Distill train epochs  : {CONFIG['distill_train_epochs']}")
print(f"   base_alpha            : {CONFIG['base_alpha']}")
print(f"   learning_rate (eta_0) : {CONFIG['learning_rate']}")
print(f"   eta_min               : {CONFIG['eta_min']}")
print(f"   lr_gamma (gamma)      : {CONFIG['lr_gamma']}")
print(f"   batch_size            : {CONFIG['batch_size']}")
print(f"   probe_batch_size      : {CONFIG['probe_batch_size']}")
print(f"   probe_grad_every      : {CONFIG['probe_gradient_every_n_epochs']} epochs")
print(f"   gradient_sample       : {CONFIG['gradient_sample_batches']} batches")
print()





# ==================================================================
# 4 OUTPUT DIRECTORY SETUP
# =======================================================================

def setup_output_directory(config):
    if IN_COLAB:
        output_root = "/content/drive/MyDrive/Math595_outputs/CIFAR10"
    elif os.path.exists("/mnt/user-data/outputs"):
        output_root = "/mnt/user-data/outputs/Math595_outputs/CIFAR10"
    elif os.path.exists("/mnt/data/outputs"):
        output_root = "/mnt/data/outputs/Math595_outputs/CIFAR10"
    else:
        output_root = "./Math595_outputs/CIFAR10"
    os.makedirs(output_root, exist_ok=True)
    run_id = config.get("_run_id") or make_run_id()
    config["_run_id"] = run_id
    run_dir = os.path.join(output_root, run_id)
    for sub in ("", "plots", "images", "arrays"):
        os.makedirs(os.path.join(run_dir, sub), exist_ok=True)
    config["output_root"] = output_root
    config["run_dir"]     = run_dir
    print(f"Output root : {output_root}")
    print(f"Run folder  : {run_dir}\n")




# 05) REPRODUCIBILITY



def set_random_seeds(seed):
    os.environ["PYTHONHASHSEED"] = str(seed)
    random.seed(seed)
    np.random.seed(seed)
    tf.random.set_seed(seed)
    print(f"Random seed set to: {seed}\n")



# 06) DATA LOADING


def load_cifar10():
    print("Loading CIFAR-10 dataset...")
    (x_train, y_train), (x_test, y_test) = tf.keras.datasets.cifar10.load_data()
    x_train = x_train.astype("float32") / 255.0
    x_test  = x_test .astype("float32") / 255.0
    y_train = y_train.flatten().astype("int64")
    y_test  = y_test .flatten().astype("int64")
    print(f"  Train : {x_train.shape}")
    print(f"  Test  : {x_test .shape}\n")
    return (x_train, y_train), (x_test, y_test)



# 007) MODEL ARCHITECTURES
# =============================================================================
def _build_baseline_cnn(config, optimizer):
    model = tf.keras.Sequential([
        tf.keras.layers.Input(shape=config["image_shape"]),
        tf.keras.layers.Conv2D(32, (3, 3), activation="relu"),
        tf.keras.layers.Conv2D(64, (3, 3), activation="relu"),
        tf.keras.layers.MaxPooling2D((2, 2)),
        tf.keras.layers.Flatten(),
        tf.keras.layers.Dense(128, activation="relu"),
        tf.keras.layers.Dense(config["num_classes"], activation="softmax"),
    ])
    model.compile(optimizer=optimizer, loss="sparse_categorical_crossentropy",
                  metrics=["accuracy"])
    return model



def _build_lenet(config, optimizer):
    model = tf.keras.Sequential([
        tf.keras.layers.Input(shape=config["image_shape"]),
        tf.keras.layers.Conv2D(6,  (5, 5), activation="tanh"),
        tf.keras.layers.AveragePooling2D((2, 2)),
        tf.keras.layers.Conv2D(16, (5, 5), activation="tanh"),
        tf.keras.layers.AveragePooling2D((2, 2)),
        tf.keras.layers.Flatten(),
        tf.keras.layers.Dense(120, activation="tanh"),
        tf.keras.layers.Dense(84,  activation="tanh"),
        tf.keras.layers.Dense(config["num_classes"], activation="softmax"),
    ])
    model.compile(optimizer=optimizer, loss="sparse_categorical_crossentropy",
                  metrics=["accuracy"])
    return model




def _resnet_basic_block(x, filters, stride=1, downsample=False):
    shortcut = x
    if downsample:
        shortcut = tf.keras.layers.Conv2D(
            filters, (1, 1), strides=stride, padding="same", use_bias=False)(shortcut)
        shortcut = tf.keras.layers.BatchNormalization()(shortcut)
    x = tf.keras.layers.Conv2D(
        filters, (3, 3), strides=stride, padding="same", use_bias=False)(x)
    x = tf.keras.layers.BatchNormalization()(x)
    x = tf.keras.layers.Activation("relu")(x)
    x = tf.keras.layers.Conv2D(
        filters, (3, 3), strides=1, padding="same", use_bias=False)(x)
    x = tf.keras.layers.BatchNormalization()(x)
    x = tf.keras.layers.Add()([x, shortcut])
    x = tf.keras.layers.Activation("relu")(x)
    return x




def _build_resnet(config, optimizer):
    """ResNet-34 adapted for CIFAR-10 (32x32). [3,4,6,3] blocks = 34 layers."""
    inputs = tf.keras.layers.Input(shape=config["image_shape"])
    x = tf.keras.layers.Conv2D(64, (3, 3), strides=1, padding="same",
                                use_bias=False)(inputs)
    x = tf.keras.layers.BatchNormalization()(x)
    x = tf.keras.layers.Activation("relu")(x)
    for filters, num_blocks, first_stride in [(64,3,1),(128,4,2),(256,6,2),(512,3,2)]:
        downsample = (first_stride != 1) or (x.shape[-1] != filters)
        x = _resnet_basic_block(x, filters, stride=first_stride, downsample=downsample)
        for _ in range(1, num_blocks):
            x = _resnet_basic_block(x, filters, stride=1, downsample=False)
    x = tf.keras.layers.GlobalAveragePooling2D()(x)
    x = tf.keras.layers.Dense(config["num_classes"], activation="softmax")(x)
    model = tf.keras.Model(inputs=inputs, outputs=x)
    model.compile(optimizer=optimizer, loss="sparse_categorical_crossentropy",
                  metrics=["accuracy"])
    return model



_MODEL_BUILDERS = {"baseline": _build_baseline_cnn, "lenet": _build_lenet,
                   "resnet": _build_resnet}

def build_model(model_name, config, optimizer):
    if model_name not in _MODEL_BUILDERS:
        raise ValueError(f"Unknown model: {model_name}")
    return _MODEL_BUILDERS[model_name](config, optimizer)






# 08 SHARED PROBE BATCH CREATION

def create_shared_probe_batch(x_data, y_data, probe_size, seed):
    rng = np.random.default_rng(seed)
    indices = rng.choice(len(x_data), size=probe_size, replace=False)
    return x_data[indices], y_data[indices], indices






# ========================================================================
# 09 COMPILED TRAINING STEP
# ========================================================================
@tf.function
def train_step_compiled(model, optimizer, loss_fn, x_batch, y_batch):
    with tf.GradientTape() as tape:
        preds = model(x_batch, training=True)
        loss = loss_fn(y_batch, preds)
    grads = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    pred_labels = tf.argmax(preds, axis=1, output_type=y_batch.dtype)
    acc = tf.reduce_mean(tf.cast(tf.equal(pred_labels, y_batch), tf.float32))
    return loss, acc





# =======================================================================
# 10) COMPILED PROBE GRADIENT
# ======================================================================


@tf.function
def compute_probe_gradient(model, loss_fn, probe_x_tf, probe_y_tf):
    with tf.GradientTape() as tape:
        tape.watch(probe_x_tf)
        preds = model(probe_x_tf, training=False)
        loss = loss_fn(probe_y_tf, preds)
    input_grad = tape.gradient(loss, probe_x_tf)
    return tf.reduce_mean(tf.abs(input_grad))





# 11) ADAPTIVE LEARNING RATE TRAINING


def train_with_adaptive_lr(model, x_train, y_train, epochs, batch_size,
                          eta_0, eta_min, gamma, probe_x, probe_y, config, phase_name):
    print(f"      Optimizer: Adam adaptive LR (eta_0={eta_0}, eta_min={eta_min}, gamma={gamma})")
    print(f"      Probe batch: {len(probe_x)} samples, every {config['probe_gradient_every_n_epochs']} epochs")
    print(f"      tf.data.Dataset + @tf.function (MAXIMUM SPEED)")

    n = len(x_train)
    steps_per_epoch = math.ceil(n / batch_size)

    x_train_tf = tf.convert_to_tensor(x_train, dtype=tf.float32)
    y_train_tf = tf.convert_to_tensor(y_train, dtype=tf.int64)
    probe_x_tf = tf.constant(probe_x, dtype=tf.float32)
    probe_y_tf = tf.constant(probe_y, dtype=tf.int64)

    model.optimizer.build(model.trainable_variables)

    dataset = tf.data.Dataset.from_tensor_slices((x_train_tf, y_train_tf))
    dataset = dataset.shuffle(buffer_size=n, seed=config["random_seed"],
                              reshuffle_each_iteration=True)
    dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)

    eta_t = eta_0
    lr_history, loss_history, acc_history = [], [], []
    first_epoch_at_min = None
    epochs_at_min = 0
    loss_fn = tf.keras.losses.SparseCategoricalCrossentropy()
    s_t = 0.0
    total_start_time = time.time()

    for epoch in range(epochs):
        model.optimizer.learning_rate.assign(eta_t)
        loss_sum, acc_sum, batch_count = 0.0, 0.0, 0

        for x_batch, y_batch in dataset:
            loss, acc = train_step_compiled(model, model.optimizer, loss_fn, x_batch, y_batch)
            loss_sum += float(loss.numpy())
            acc_sum += float(acc.numpy())
            batch_count += 1
            if batch_count % 200 == 0:
                print(f"      Epoch {epoch+1}/{epochs} - Batch {batch_count}/{steps_per_epoch} "
                      f"[{batch_count/steps_per_epoch*100:.0f}%]", end='\r')

        avg_loss = loss_sum / batch_count
        avg_acc = acc_sum / batch_count
        loss_history.append(avg_loss)
        acc_history.append(avg_acc)
        lr_history.append((epoch + 1, eta_t, s_t))

        if (epoch + 1) % config["probe_gradient_every_n_epochs"] == 0 or epoch == 0:
            s_t = float(compute_probe_gradient(model, loss_fn, probe_x_tf, probe_y_tf).numpy())
            eta_next = max(eta_min, eta_t - gamma * eta_0 * s_t)
            if eta_next == eta_min:
                epochs_at_min += 1
                if first_epoch_at_min is None:
                    first_epoch_at_min = epoch + 1
            eta_t = eta_next

        if epoch == 0:
            t1 = time.time() - total_start_time
            print(f"\n      Estimated time: {t1*(epochs-1)/60:.1f} min ({t1:.1f}s/epoch)")
            print(f"      Note: First epoch includes compilation overhead")

        if (epoch + 1) % 10 == 0 or (epoch + 1) == epochs:
            elapsed = time.time() - total_start_time
            avg_t = elapsed / (epoch + 1)
            rem = avg_t * (epochs - epoch - 1)
            print(f"      Epoch {epoch+1}/{epochs}  loss={avg_loss:.4f}  acc={avg_acc:.4f}  "
                  f"eta={eta_t:.6e}  s={s_t:.6e}  {avg_t:.1f}s/ep  ETA:{rem/60:.1f}min")

    total_time = time.time() - total_start_time
    print(f"      Done in {total_time:.1f}s ({total_time/epochs:.1f}s/epoch avg)")

    return ({"loss": loss_history, "accuracy": acc_history},
            lr_history,
            {"first_epoch_at_min": first_epoch_at_min,
             "epochs_at_min": epochs_at_min,
             "percent_at_min": (epochs_at_min / epochs) * 100 if epochs > 0 else 0,
             "total_training_time_s": total_time,
             "avg_time_per_epoch_s": total_time / epochs})




# 12 GLOBAL WEIGHT-GRADIENT EXTRACTION
# ===========================================================================

def extract_global_weight_gradient(model, x_train, y_train, batch_size, sample_batches=None):
    n = len(x_train)
    num_batches = math.ceil(n / batch_size)
    x_tf = tf.convert_to_tensor(x_train, dtype=tf.float32)
    y_tf = tf.convert_to_tensor(y_train, dtype=tf.int64)

    if sample_batches and sample_batches < num_batches:
        rng = np.random.default_rng(42)
        batch_indices = sorted(rng.choice(num_batches, size=sample_batches, replace=False))
        print(f"    Gradient extraction: sampling {sample_batches}/{num_batches} batches")
    else:
        batch_indices = range(num_batches)
        print(f"    Gradient extraction: {n} images, {num_batches} batches")

    batch_norms = []
    t0 = time.time()
    for idx, b in enumerate(batch_indices):
        s, e = b * batch_size, min((b + 1) * batch_size, n)
        with tf.GradientTape() as tape:
            preds = model(x_tf[s:e], training=False)
            loss = tf.reduce_mean(
                tf.keras.losses.sparse_categorical_crossentropy(y_tf[s:e], preds))
        grads = tape.gradient(loss, model.trainable_variables)
        l1 = tf.add_n([tf.reduce_sum(tf.abs(g)) for g in grads if g is not None])
        batch_norms.append(float(l1.numpy()))
        if (idx + 1) % 10 == 0 or (idx + 1) == len(batch_indices):
            print(f"      {idx+1}/{len(batch_indices)} ({(idx+1)/len(batch_indices)*100:.0f}%)",
                  end="\r" if (idx+1) < len(batch_indices) else "\n")

    elapsed = time.time() - t0
    bn = np.array(batch_norms, dtype=np.float64)
    g_global = float(bn.mean())
    g_global_norm = g_global / (float(bn.max()) + 1e-12)

    print(f"      Done ({elapsed:.1f}s)  g_global_norm={g_global_norm:.6e}\n")
    return g_global_norm, {
        "batch_norms_min": float(bn.min()), "batch_norms_mean": g_global,
        "batch_norms_max": float(bn.max()), "g_global": g_global,
        "g_global_norm": g_global_norm, "extraction_time_s": elapsed,
        "batches_sampled": len(batch_indices), "batches_total": num_batches}




# =============================================================================
# 13 IMAGE MODIFICATION
# =============================================================================


def modify_images_constant(x_train, base_alpha, g_global_norm):
    delta = base_alpha * g_global_norm
    x_mod = np.clip(x_train - delta, 0.0, 1.0).astype(np.float32)
    print(f"      delta = {delta:.6e}  (alpha={base_alpha:.1e} x g={g_global_norm:.6e})")
    if delta > 0.1:
        print("      WARNING: delta > 0.1 -- images heavily modified!")
    elif delta < 1e-5:
        print("      WARNING: delta < 1e-5 -- images barely modified!")
    return x_mod




# 14. CLASS-WISE K-MEANS CLUSTERING


def cluster_classwise(x_modified, y_train, k_per_class, config):
    nc = config["num_classes"]
    img_h, img_w, img_c = config["image_shape"]
    distilled_images, distilled_labels = [], []
    print(f"      KMeans: {k_per_class} clusters/class x {nc} classes = {k_per_class*nc} images")
    t0 = time.time()
    for c in range(nc):
        idx = np.where(y_train == c)[0]
        x_flat = x_modified[idx].reshape(len(idx), -1)
        km = MiniBatchKMeans(n_clusters=k_per_class, n_init=config["kmeans_n_init"],
                             max_iter=config["kmeans_max_iter"],
                             batch_size=config["kmeans_batch_size"],
                             random_state=config["random_seed"])
        km.fit(x_flat)
        centroids = np.clip(km.cluster_centers_.reshape(k_per_class, img_h, img_w, img_c), 0, 1)
        distilled_images.append(centroids)
        distilled_labels.extend([c] * k_per_class)
        print(f"        class {c} ({CIFAR10_CLASSES[c]}) done", end="\r")
    print(f"      Clustering complete ({time.time()-t0:.1f}s)\n")
    return (np.concatenate(distilled_images, axis=0).astype(np.float32),
            np.array(distilled_labels, dtype=np.int64))




# 15> DIAGNOSTIC METRICS


def compute_diagnostics(model, x_test, y_test, class_names):
    """Compute accuracy, sensitivity, specificity, F1, confusion matrix."""
    y_pred = np.argmax(model.predict(x_test, verbose=0), axis=1)
    y_true = y_test
    nc = len(class_names)

    cm_raw = confusion_matrix(y_true, y_pred, labels=list(range(nc)))
    row_sums = cm_raw.sum(axis=1, keepdims=True).astype("float64")
    row_sums[row_sums == 0] = 1
    cm_norm = cm_raw.astype("float64") / row_sums

    per_class_recall = np.diag(cm_raw).astype("float64")
    support = cm_raw.sum(axis=1).astype("float64")
    support[support == 0] = 1
    per_class_recall = per_class_recall / support

    per_class_spec = np.zeros(nc, dtype="float64")
    for c in range(nc):
        tp = cm_raw[c, c]
        fn = cm_raw[c, :].sum() - tp
        fp = cm_raw[:, c].sum() - tp
        tn = cm_raw.sum() - tp - fn - fp
        per_class_spec[c] = tn / (tn + fp) if (tn + fp) > 0 else 0.0

    prec, rec, f1_per, sup = precision_recall_fscore_support(
        y_true, y_pred, labels=list(range(nc)), zero_division=0)

    return {
        "accuracy":              float(np.mean(y_pred == y_true)),
        "sensitivity_macro":     float(per_class_recall.mean()),
        "specificity_macro":     float(per_class_spec.mean()),
        "macro_f1":              float(f1_score(y_true, y_pred, average="macro", zero_division=0)),
        "per_class_recall":      per_class_recall,
        "per_class_specificity": per_class_spec,
        "per_class_precision":   prec,
        "per_class_f1":          f1_per,
        "per_class_support":     sup,
        "cm_raw":                cm_raw,
        "cm_normalized":         cm_norm,
        "classification_report": classification_report(
            y_true, y_pred, labels=list(range(nc)),
            target_names=class_names, digits=4, zero_division=0),
    }




def print_diagnostics(diag, phase_label, class_names):
    w = 64
    print(f"\n      +{'-'*w}+")
    print(f"      |  DIAGNOSTICS: {phase_label:<{w-16}}|")
    print(f"      +{'-'*w}+")
    print(f"      |  Accuracy           : {diag['accuracy']*100:>7.2f} %{' '*(w-36)}|")
    print(f"      |  Sensitivity (macro) : {diag['sensitivity_macro']*100:>7.2f} %{' '*(w-36)}|")
    print(f"      |  Specificity (macro) : {diag['specificity_macro']*100:>7.2f} %{' '*(w-36)}|")
    print(f"      |  Macro-F1            : {diag['macro_f1']*100:>7.2f} %{' '*(w-36)}|")
    print(f"      +{'-'*w}+")
    print(f"      |  {'Class':<12} {'Recall':>8} {'Specif.':>8} {'Prec.':>8} {'F1':>8}{'':>{w-52}}|")
    print(f"      |  {'---'*18:<{w-4}}|")
    for i, name in enumerate(class_names):
        print(f"      |  {name:<12} "
              f"{diag['per_class_recall'][i]*100:>7.2f}% "
              f"{diag['per_class_specificity'][i]*100:>7.2f}% "
              f"{diag['per_class_precision'][i]*100:>7.2f}% "
              f"{diag['per_class_f1'][i]*100:>7.2f}%"
              f"{'':>{w-52}}|")
    print(f"      +{'-'*w}+")




def save_confusion_matrix_plot(cm_norm, class_names, title, out_path):
    fig, ax = plt.subplots(figsize=(10, 8))
    im = ax.imshow(cm_norm, interpolation="nearest", cmap="Blues", vmin=0, vmax=1)
    fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
    ax.set_xticks(range(len(class_names)))
    ax.set_yticks(range(len(class_names)))
    ax.set_xticklabels(class_names, rotation=45, ha="right", fontsize=9)
    ax.set_yticklabels(class_names, fontsize=9)
    for i in range(len(class_names)):
        for j in range(len(class_names)):
            v = cm_norm[i, j]
            ax.text(j, i, f"{v:.2f}", ha="center", va="center",
                    color="white" if v > 0.5 else "black", fontsize=8)
    ax.set_xlabel("Predicted", fontsize=11)
    ax.set_ylabel("True", fontsize=11)
    ax.set_title(title, fontsize=13)
    plt.tight_layout()
    plt.savefig(out_path, dpi=150, bbox_inches="tight")
    plt.close(fig)




# ===================================================================================
# 16.... PLOT & GRID HELPERS


def _save_curve(history, title, key_train, key_val, ylabel, out_path):
    if history is None: return
    h = history if isinstance(history, dict) else history.history
    plt.figure(figsize=(8, 4))
    if key_train in h: plt.plot(h[key_train], label="train")
    if key_val in h:   plt.plot(h[key_val],   label="val")
    plt.title(title); plt.xlabel("Epoch"); plt.ylabel(ylabel); plt.legend()
    plt.tight_layout(); plt.savefig(out_path, dpi=100, bbox_inches="tight"); plt.close()

def save_history_plot(h, title, p):
    _save_curve(h, title+" - Loss", "loss", "val_loss", "Loss", p)
def save_acc_plot(h, title, p):
    _save_curve(h, title+" - Accuracy", "accuracy", "val_accuracy", "Accuracy", p)

def save_distilled_grid(x_dist, y_dist, k, out_path):
    nc = 10; order = np.argsort(y_dist); x = x_dist[order]
    col_w = max(0.4, 8.0/k); fig_w = col_w*k; fig_h = 0.6*nc
    fig, axes = plt.subplots(nc, k, figsize=(fig_w, fig_h), squeeze=False)
    for r in range(nc):
        for c in range(k):
            axes[r,c].imshow(np.clip(x[r*k+c], 0, 1)); axes[r,c].axis("off")
    fig.suptitle(f"Distilled Images (K={k})", fontsize=11, y=1.02)
    plt.tight_layout(); plt.savefig(out_path, dpi=100, bbox_inches="tight"); plt.close(fig)








# 17...... ONE FULL EXPERIMENT
# ========================================================================================

def run_one_experiment(model_name, config, x_train, y_train, x_test, y_test,
                       shared_probe_x, shared_probe_y, shared_probe_indices):
    tag = f"{model_name}_ADAPTIVE"
    dname = {"baseline": "Baseline CNN", "lenet": "LeNet", "resnet": "ResNet"}[model_name]

    print(f"\n{'='*80}")
    print(f"  {dname}  x  ADAPTIVE LR")
    print(f"{'='*80}")
    t_exp = time.time()



    # -- A) TRAIN ON ORIGINAL (FULL) DATA ------------------------------------
    print(f"\n  [A] Train on original data ({config['full_train_epochs']} epochs)")
    opt = tf.keras.optimizers.Adam(learning_rate=config["learning_rate"])
    orig_model = build_model(model_name, config, opt)
    t0 = time.time()
    orig_hist, orig_lr_hist, orig_lr_stats = train_with_adaptive_lr(
        orig_model, x_train, y_train, config["full_train_epochs"], config["batch_size"],
        config["learning_rate"], config["eta_min"], config["lr_gamma"],
        shared_probe_x, shared_probe_y, config, "original")
    orig_time = time.time() - t0

    print(f"\n  [A.eval] Original model diagnostics")
    orig_diag = compute_diagnostics(orig_model, x_test, y_test, CIFAR10_CLASSES)
    print_diagnostics(orig_diag, f"{dname} -- Original (Full) Data", CIFAR10_CLASSES)


    # -- B+C+D) DISTILLATION (gradient + modify + KMeans) ------------
    #    T_distill = total time for the distillation process
    t_distill_start = time.time()


    # -- B) GRADIENT SIGNAL ---------------------------------------
    print(f"\n  [B] Weight gradient extraction")
    g_norm, g_stats = extract_global_weight_gradient(
        orig_model, x_train, y_train, config["gradient_batch_size"],
        sample_batches=config["gradient_sample_batches"])



    # -- C) MODIFY IMAGES ----------------------------------
    print(f"  [C] Image modification")
    x_mod = modify_images_constant(x_train, config["base_alpha"], g_norm)



    # -- D) K-MEANS ------------------------------------------------
    print(f"\n  [D] Class-wise KMeans (K={config['k_per_class']})")
    x_dist, y_dist = cluster_classwise(x_mod, y_train, config["k_per_class"], config)

    distill_time = time.time() - t_distill_start
    print(f"      T_distill (B+C+D total): {distill_time:.1f} s\n")




    # -- E) TRAIN ON DISTILLED DATA --------------------------------
    print(f"  [E] Train on {len(x_dist)} distilled images ({config['distill_train_epochs']} epochs)")
    dist_probe_sz = min(config["probe_batch_size"], len(x_dist))
    dist_probe_x, dist_probe_y, _ = create_shared_probe_batch(
        x_dist, y_dist, dist_probe_sz, config["random_seed"])
    print(f"      Distilled probe: {dist_probe_sz} samples from distilled data")





    opt2 = tf.keras.optimizers.Adam(learning_rate=config["learning_rate"])
    dist_model = build_model(model_name, config, opt2)
    t0 = time.time()
    dist_hist, dist_lr_hist, dist_lr_stats = train_with_adaptive_lr(
        dist_model, x_dist, y_dist, config["distill_train_epochs"], config["batch_size"],
        config["learning_rate"], config["eta_min"], config["lr_gamma"],
        dist_probe_x, dist_probe_y, config, "distilled")
    



    
    dist_time = time.time() - t0




    print(f"\n  [E.eval] Distilled model diagnostics")
    dist_diag = compute_diagnostics(dist_model, x_test, y_test, CIFAR10_CLASSES)
    print_diagnostics(dist_diag, f"{dname} -- Distilled Data", CIFAR10_CLASSES)



    # -- F) SUMMARY BOX ----------------------
    drop = orig_diag["accuracy"] - dist_diag["accuracy"]
    total_wall = time.time() - t_exp




    print(f"\n      +{'='*64}+")
    print(f"      |  {dname} -- SUMMARY{' '*(64-len(dname)-14)}|")
    print(f"      +{'='*64}+")
    print(f"      |  Full Train Accuracy        : {orig_diag['accuracy']*100:>7.2f} %{' '*25}|")
    print(f"      |  Distill Train Accuracy      : {dist_diag['accuracy']*100:>7.2f} %{' '*25}|")
    print(f"      |  Drop                       : {drop*100:>7.2f} %{' '*25}|")
    print(f"      +{'-'*64}+")
    print(f"      |  Orig Sensitivity (macro)   : {orig_diag['sensitivity_macro']*100:>7.2f} %{' '*25}|")
    print(f"      |  Dist Sensitivity (macro)   : {dist_diag['sensitivity_macro']*100:>7.2f} %{' '*25}|")
    print(f"      |  Orig Specificity (macro)   : {orig_diag['specificity_macro']*100:>7.2f} %{' '*25}|")
    print(f"      |  Dist Specificity (macro)   : {dist_diag['specificity_macro']*100:>7.2f} %{' '*25}|")
    print(f"      |  Orig Macro-F1              : {orig_diag['macro_f1']*100:>7.2f} %{' '*25}|")
    print(f"      |  Dist Macro-F1              : {dist_diag['macro_f1']*100:>7.2f} %{' '*25}|")
    print(f"      +{'-'*64}+")
    print(f"      |  T_full_train  (full data)  : {orig_time:>9.1f} s{' '*23}|")
    print(f"      |  T_distill     (B+C+D)      : {distill_time:>9.1f} s{' '*23}|")
    print(f"      |  T_distill_train (distilled) : {dist_time:>9.1f} s{' '*23}|")
    print(f"      |  Total wall time            : {total_wall:>9.1f} s{' '*23}|")
    print(f"      +{'='*64}+\n")

    del orig_model, dist_model
    tf.keras.backend.clear_session()
    print("      [cleanup] done\n")



    return {
        "Dataset": config["dataset"], "Model": dname, "K": config["k_per_class"],
        "FullTrainEpochs": config["full_train_epochs"],
        "DistillTrainEpochs": config["distill_train_epochs"],
        "AlphaType": "CONSTANT", "LRMode": "ADAPTIVE",
        "FullTrainAcc": orig_diag["accuracy"],
        "DistillTrainAcc": dist_diag["accuracy"], "Drop": drop,
        "FullTrainTime_s": orig_time,
        "DistillTime_s": distill_time,
        "DistillTrainTime_s": dist_time,
        "_original_diag": orig_diag, "_distilled_diag": dist_diag,
        "_tag": tag, "_x_distilled": x_dist, "_y_distilled": y_dist,
        "_original_history": orig_hist, "_distilled_history": dist_hist,
        "_original_lr_history": orig_lr_hist, "_distilled_lr_history": dist_lr_hist,
        "_original_lr_stats": orig_lr_stats, "_distilled_lr_stats": dist_lr_stats,
        "_g_stats": g_stats,
    }





# =============================================================================
# 18) SUMMARY TEXT
# =============================================================================

def _fmt_diag_txt(diag, label, cnames):
    o = StringIO()
    o.write(f"\n  {label}\n  {'='*70}\n")
    o.write(f"  Accuracy            : {diag['accuracy']*100:.2f}%\n")
    o.write(f"  Sensitivity (macro) : {diag['sensitivity_macro']*100:.2f}%\n")
    o.write(f"  Specificity (macro) : {diag['specificity_macro']*100:.2f}%\n")
    o.write(f"  Macro-F1            : {diag['macro_f1']*100:.2f}%\n\n")
    o.write(f"  {'Class':<12} {'Recall':>8} {'Specif':>8} {'Prec':>8} {'F1':>8} {'Support':>8}\n")
    o.write(f"  {'-'*56}\n")
    for i, nm in enumerate(cnames):
        o.write(f"  {nm:<12} {diag['per_class_recall'][i]*100:>7.2f}% "
                f"{diag['per_class_specificity'][i]*100:>7.2f}% "
                f"{diag['per_class_precision'][i]*100:>7.2f}% "
                f"{diag['per_class_f1'][i]*100:>7.2f}% "
                f"{int(diag['per_class_support'][i]):>8}\n")
    o.write(f"\n  sklearn classification_report:\n")
    for line in diag["classification_report"].split("\n"):
        o.write(f"  {line}\n")
    return o.getvalue()



def generate_summary(all_results, config, probe_indices):
    o = StringIO()
    o.write("=" * 80 + "\n")
    o.write("CIFAR-10 DISTILLATION -- ADAPTIVE LEARNING RATE \n")
    o.write("=" * 80 + "\n")
    o.write(f"Run ID                : {config.get('_run_id','?')}\n")
    o.write(f"Date (CST)            : {now_central().strftime('%B %d, %Y, %I:%M %p %Z')}\n")
    o.write(f"Dataset               : {config['dataset']}\n")
    o.write(f"K per class           : {config['k_per_class']}\n")
    o.write(f"Total distilled imgs  : {config['k_per_class']*config['num_classes']}\n")
    o.write(f"Full train epochs     : {config['full_train_epochs']}\n")
    o.write(f"Distill train epochs  : {config['distill_train_epochs']}\n")
    o.write(f"base_alpha            : {config['base_alpha']}\n")
    o.write(f"learning_rate (eta_0) : {config['learning_rate']}\n")
    o.write(f"eta_min               : {config['eta_min']}\n")
    o.write(f"lr_gamma (gamma)      : {config['lr_gamma']}\n")
    o.write(f"batch_size            : {config['batch_size']}\n")
    o.write(f"probe_batch_size      : {config['probe_batch_size']}\n")
    o.write(f"probe_grad_every      : {config['probe_gradient_every_n_epochs']} epochs\n")
    o.write(f"probe_indices         : {probe_indices.tolist()}\n")
    o.write(f"gradient_sample       : {config['gradient_sample_batches']} batches\n")
    o.write("=" * 80 + "\n\n")




    hdr = (f"{'Model':<14} {'FullAcc':>8} {'DistAcc':>8} {'Drop':>7} "
           f"{'DistSens':>9} {'DistSpec':>9} {'DistF1':>7} "
           f"{'T_full':>7} {'T_dist':>7} {'T_dTrn':>7}")
    o.write(hdr + "\n" + "-"*len(hdr) + "\n")
    for r in all_results:
        dd = r.get("_distilled_diag", {})
        o.write(f"{r['Model']:<14} {r['FullTrainAcc']*100:>7.2f}% "
                f"{r['DistillTrainAcc']*100:>7.2f}% {r['Drop']*100:>6.2f}% "
                f"{dd.get('sensitivity_macro',0)*100:>8.2f}% "
                f"{dd.get('specificity_macro',0)*100:>8.2f}% "
                f"{dd.get('macro_f1',0)*100:>6.2f}% "
                f"{r.get('FullTrainTime_s',0):>6.1f}s "
                f"{r.get('DistillTime_s',0):>6.1f}s "
                f"{r.get('DistillTrainTime_s',0):>6.1f}s\n")
    o.write("\n" + "="*80 + "\n")



    for r in all_results:
        o.write(f"\n{'='*80}\n  {r['Model']} -- FULL DIAGNOSTICS\n{'='*80}\n")
        o.write(f"\n  Runtime: T_full_train={r.get('FullTrainTime_s',0):.1f}s  "
                f"T_distill={r.get('DistillTime_s',0):.1f}s  "
                f"T_distill_train={r.get('DistillTrainTime_s',0):.1f}s\n")
        if r.get("_original_diag"):
            o.write(_fmt_diag_txt(r["_original_diag"],
                    f"{r['Model']} -- Original (Full) Data", CIFAR10_CLASSES))
        if r.get("_distilled_diag"):
            o.write(_fmt_diag_txt(r["_distilled_diag"],
                    f"{r['Model']} -- Distilled Data", CIFAR10_CLASSES))

        for phase, lrk, stk in [("ORIGINAL", "_original_lr_history", "_original_lr_stats"),
                                 ("DISTILLED", "_distilled_lr_history", "_distilled_lr_stats")]:
            st = r.get(stk, {})
            o.write(f"\n  {r['Model']} -- {phase} MODEL LR EVOLUTION\n  {'-'*60}\n")
            o.write(f"  First eta_min: epoch {st.get('first_epoch_at_min','never')}, "
                    f"{st.get('percent_at_min',0):.1f}% at eta_min\n")
            o.write(f"  Time: {st.get('total_training_time_s',0):.1f}s "
                    f"({st.get('avg_time_per_epoch_s',0):.2f}s/ep)\n")
            o.write(f"  {'Epoch':>6} {'eta':>12} {'s_t':>12}\n  {'-'*32}\n")
            lh = r.get(lrk, [])
            for ep, et, st_val in lh[:10]:
                o.write(f"  {ep:>6} {et:>12.6e} {st_val:>12.6e}\n")
            if len(lh) > 20:
                o.write("  ...\n")
                for ep, et, st_val in lh[-10:]:
                    o.write(f"  {ep:>6} {et:>12.6e} {st_val:>12.6e}\n")

    o.write("\n" + "="*80 + "\n")
    if all_results:
        best = max(all_results, key=lambda x: x["DistillTrainAcc"])
        o.write(f"BEST DISTILLED MODEL: {best['Model']} | "
                f"{best['DistillTrainAcc']*100:.2f}%\n")
    o.write("="*80 + "\n")
    return o.getvalue()






# ======================================================================
# 19) SAVE ALL OUTPUTS
# ======================================================================

def save_all(all_results, config, probe_indices):
    rd = config["run_dir"]
    pd = os.path.join(rd, "plots")
    imd = os.path.join(rd, "images")
    ard = os.path.join(rd, "arrays")

    # CSV
    fields = ["Dataset","Model","K","FullTrainEpochs","DistillTrainEpochs",
              "AlphaType","LRMode",
              "FullTrainAcc","DistillTrainAcc","Drop",
              "OrigSensitivity","OrigSpecificity","OrigMacroF1",
              "DistSensitivity","DistSpecificity","DistMacroF1",
              "FullTrainTime_s","DistillTime_s","DistillTrainTime_s"]
    cp = os.path.join(rd, "results.csv")
    with open(cp, "w", newline="", encoding="utf-8") as fh:
        w = csv.DictWriter(fh, fieldnames=fields); w.writeheader()
        for r in all_results:
            od = r.get("_original_diag", {}); dd = r.get("_distilled_diag", {})
            w.writerow({
                "Dataset": r["Dataset"], "Model": r["Model"], "K": r["K"],
                "FullTrainEpochs": r["FullTrainEpochs"],
                "DistillTrainEpochs": r["DistillTrainEpochs"],
                "AlphaType": r["AlphaType"], "LRMode": r["LRMode"],
                "FullTrainAcc": f"{r['FullTrainAcc']*100:.2f}%",
                "DistillTrainAcc": f"{r['DistillTrainAcc']*100:.2f}%",
                "Drop": f"{r['Drop']*100:.2f}%",
                "OrigSensitivity": f"{od.get('sensitivity_macro',0)*100:.2f}%",
                "OrigSpecificity": f"{od.get('specificity_macro',0)*100:.2f}%",
                "OrigMacroF1": f"{od.get('macro_f1',0)*100:.2f}%",
                "DistSensitivity": f"{dd.get('sensitivity_macro',0)*100:.2f}%",
                "DistSpecificity": f"{dd.get('specificity_macro',0)*100:.2f}%",
                "DistMacroF1": f"{dd.get('macro_f1',0)*100:.2f}%",
                "FullTrainTime_s": f"{r.get('FullTrainTime_s',0):.1f}",
                "DistillTime_s": f"{r.get('DistillTime_s',0):.1f}",
                "DistillTrainTime_s": f"{r.get('DistillTrainTime_s',0):.1f}",
            })
    print(f"  CSV  : {cp}")


    # Summary
    tp = os.path.join(rd, "summary.txt")
    with open(tp, "w", encoding="utf-8") as fh:
        fh.write(generate_summary(all_results, config, probe_indices))
    print(f"  TXT  : {tp}")




    # Config
    jp = os.path.join(rd, "config.json")
    with open(jp, "w", encoding="utf-8") as fh:
        json.dump(config, fh, indent=2, default=str)
    print(f"  JSON : {jp}")



    # Per-experiment artefacts
    for r in all_results:
        tag = r["_tag"]
        if config.get("save_plots"):
            for phase, hk in [("original","_original_history"),("distilled","_distilled_history")]:
                h = r.get(hk)
                if h is None: continue
                try:
                    save_history_plot(h, f"{phase.title()} ({tag})",
                                     os.path.join(pd, f"{phase}_loss_{tag}.png"))
                    save_acc_plot(h, f"{phase.title()} ({tag})",
                                 os.path.join(pd, f"{phase}_acc_{tag}.png"))
                except Exception as e:
                    print(f"  [WARN] plots {tag}: {e}")

            for phase, dk in [("original","_original_diag"),("distilled","_distilled_diag")]:
                d = r.get(dk)
                if d is None: continue
                cmp = os.path.join(pd, f"confusion_matrix_{phase}_{tag}.png")
                try:
                    save_confusion_matrix_plot(d["cm_normalized"], CIFAR10_CLASSES,
                        f"CM - {phase.title()} ({r['Model']})", cmp)
                    print(f"  CM   : {cmp}")
                except Exception as e:
                    print(f"  [WARN] CM {tag}: {e}")

        if config.get("save_distilled_grids") and r.get("_x_distilled") is not None:
            gp = os.path.join(imd, f"distilled_grid_{tag}_K{r['K']}.png")
            try:
                save_distilled_grid(r["_x_distilled"], r["_y_distilled"], r["K"], gp)
                print(f"  GRID : {gp}")
            except Exception as e:
                print(f"  [WARN] grid {tag}: {e}")

        if config.get("save_arrays") and r.get("_x_distilled") is not None:
            for arr, nm in [(r["_x_distilled"], "x"), (r["_y_distilled"], "y")]:
                ap = os.path.join(ard, f"{nm}_distilled_{tag}_K{r['K']}.npy")
                try:
                    np.save(ap, arr); print(f"  NPY  : {ap}")
                except Exception as e:
                    print(f"  [WARN] {nm} {tag}: {e}")




# 20) MAIN

def main():
    CONFIG["_run_id"] = make_run_id()
    setup_output_directory(CONFIG)
    set_random_seeds(CONFIG["random_seed"])

    (x_train, y_train), (x_test, y_test) = load_cifar10()

    print("Creating shared probe batch (original data)...")
    shared_px, shared_py, shared_idx = create_shared_probe_batch(
        x_train, y_train, CONFIG["probe_batch_size"], CONFIG["random_seed"])
    print(f"Probe indices: {shared_idx[:10]}... ({len(shared_idx)} total)\n")



    print("=" * 80)
    print("  CIFAR-10 DISTILLATION -- ADAPTIVE LR (SIRAKOV METHOD)")
    print("=" * 80)
    print(f"  Run folder           : {CONFIG['run_dir']}")
    print(f"  Date (CST)           : {now_central().strftime('%B %d, %Y, %I:%M %p %Z')}")
    print(f"  K per class          : {CONFIG['k_per_class']}")
    print(f"  Full train epochs    : {CONFIG['full_train_epochs']}")
    print(f"  Distill train epochs : {CONFIG['distill_train_epochs']}")
    print(f"  base_alpha           : {CONFIG['base_alpha']}")
    print(f"  eta_0                : {CONFIG['learning_rate']}")
    print(f"  eta_min              : {CONFIG['eta_min']}")
    print(f"  gamma                : {CONFIG['lr_gamma']}")
    print(f"  batch_size           : {CONFIG['batch_size']}")
    print("=" * 80 + "\n")


    models = ["baseline", "lenet", "resnet"]
    _nm = {"baseline": "Baseline CNN", "lenet": "LeNet", "resnet": "ResNet"}
    all_results = []

    for mn in models:
        print(f"\n{'='*80}\n    MODEL: {_nm[mn]}\n{'='*80}")
        result = run_one_experiment(mn, CONFIG, x_train, y_train, x_test, y_test,
                                   shared_px, shared_py, shared_idx)
        all_results.append(result)

    print("\n" + "="*80 + "\n  SAVING OUTPUTS\n" + "="*80)
    save_all(all_results, CONFIG, shared_idx)
    print("\n" + generate_summary(all_results, CONFIG, shared_idx))
    print("=" * 80 + "\n  COMPLETE\n" + "=" * 80)




if __name__ == "__main__":
    main()