"""
Petals to the Metal: Flower Classification on TPU
State-of-the-Art EfficientNet Deep Learning Pipeline
Author: Matheus Bonjour
Task: 104 Flower Species Classification (Macro F1 Metric)
"""

import os
import sys
import re
import gc
import glob
import math
import numpy as np
import pandas as pd
import tensorflow as tf

print(f"=== Running Petals to the Metal Flower Classification | TensorFlow Version: {tf.__version__} ===")

# ==============================================================================
# Hardware Strategy (TPU / GPU / CPU Auto-Detection)
# ==============================================================================
try:
    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()
    print('Running on TPU ', tpu.master())
    tf.config.experimental_connect_to_cluster(tpu)
    tf.tpu.experimental.initialize_tpu_system(tpu)
    strategy = tf.distribute.TPUStrategy(tpu)
except Exception as e:
    print(f"TPU initialization failed ({e}). Checking GPU/CPU strategy...")
    strategy = tf.distribute.get_strategy()

print("REPLICAS: ", strategy.num_replicas_in_sync)

# ==============================================================================
# Configuration Parameters
# ==============================================================================
IMAGE_SIZE = [512, 512]
BATCH_SIZE = 16 * strategy.num_replicas_in_sync
EPOCHS = 16
NUM_CLASSES = 104

def find_file(filename):
    search_paths = ["/kaggle/input", "."]
    for sp in search_paths:
        matches = glob.glob(os.path.join(sp, "**", filename), recursive=True)
        if matches:
            print(f"Found {filename} at: {matches[0]}")
            return matches[0]
    raise FileNotFoundError(f"Could not locate {filename}")

def get_data_dir():
    search_paths = [
        "/kaggle/input/tpu-getting-started",
        "/kaggle/input/competitions/tpu-getting-started",
        "/kaggle/input"
    ]
    for sp in search_paths:
        if os.path.exists(sp):
            tfrec_dirs = glob.glob(os.path.join(sp, "**", "tfrecords-jpeg-512x512"), recursive=True)
            if tfrec_dirs:
                return tfrec_dirs[0]
    # Fallback to any tfrecords directory
    tfrec_dirs = glob.glob(os.path.join("/kaggle/input", "**", "*.tfrec"), recursive=True)
    if tfrec_dirs:
        return os.path.dirname(os.path.dirname(tfrec_dirs[0]))
    return "."

DATA_DIR = get_data_dir()
print(f"Data Directory: {DATA_DIR}")

TRAIN_FILES = tf.io.gfile.glob(os.path.join(DATA_DIR, "train", "*.tfrec"))
VAL_FILES = tf.io.gfile.glob(os.path.join(DATA_DIR, "val", "*.tfrec"))
TEST_FILES = tf.io.gfile.glob(os.path.join(DATA_DIR, "test", "*.tfrec"))

print(f"Train TFRecord files: {len(TRAIN_FILES)}")
print(f"Val TFRecord files: {len(VAL_FILES)}")
print(f"Test TFRecord files: {len(TEST_FILES)}")

# ==============================================================================
# TFRecord Parsing & Augmentation API
# ==============================================================================
def decode_image(image_data):
    image = tf.image.decode_jpeg(image_data, channels=3)
    image = tf.cast(image, tf.float32) / 255.0
    image = tf.reshape(image, [*IMAGE_SIZE, 3])
    return image

def read_labeled_tfrecord(example):
    LABELED_TFREC_FORMAT = {
        "image": tf.io.FixedLenFeature([], tf.string),
        "class": tf.io.FixedLenFeature([], tf.int64),
    }
    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)
    image = decode_image(example['image'])
    label = tf.cast(example['class'], tf.int32)
    return image, label

def read_unlabeled_tfrecord(example):
    UNLABELED_TFREC_FORMAT = {
        "image": tf.io.FixedLenFeature([], tf.string),
        "id": tf.io.FixedLenFeature([], tf.string),
    }
    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)
    image = decode_image(example['image'])
    idnum = example['id']
    return image, idnum

def data_augment(image, label):
    image = tf.image.random_flip_left_right(image)
    image = tf.image.random_flip_up_down(image)
    image = tf.image.random_saturation(image, 0.8, 1.2)
    image = tf.image.random_brightness(image, 0.1)
    return image, label

def load_dataset(filenames, labeled=True, ordered=False):
    ignore_order = tf.data.Options()
    if not ordered:
        ignore_order.experimental_deterministic = False

    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=tf.data.AUTOTUNE)
    dataset = dataset.with_options(ignore_order)
    dataset = dataset.map(read_labeled_tfrecord if labeled else read_unlabeled_tfrecord, num_parallel_calls=tf.data.AUTOTUNE)
    return dataset

def get_training_dataset():
    dataset = load_dataset(TRAIN_FILES + VAL_FILES, labeled=True)
    dataset = dataset.map(data_augment, num_parallel_calls=tf.data.AUTOTUNE)
    dataset = dataset.repeat()
    dataset = dataset.shuffle(2048)
    dataset = dataset.batch(BATCH_SIZE)
    dataset = dataset.prefetch(tf.data.AUTOTUNE)
    return dataset

def get_test_dataset(ordered=True):
    dataset = load_dataset(TEST_FILES, labeled=False, ordered=ordered)
    dataset = dataset.batch(BATCH_SIZE)
    dataset = dataset.prefetch(tf.data.AUTOTUNE)
    return dataset

def count_data_items(filenames):
    n = [int(re.compile(r"-([0-9]+)\.").search(filename).group(1)) for filename in filenames]
    return sum(n)

NUM_TRAIN_IMAGES = count_data_items(TRAIN_FILES + VAL_FILES) if TRAIN_FILES else 12000
NUM_TEST_IMAGES = count_data_items(TEST_FILES) if TEST_FILES else 7000
STEPS_PER_EPOCH = NUM_TRAIN_IMAGES // BATCH_SIZE

print(f"Dataset summary: {NUM_TRAIN_IMAGES} training images | {NUM_TEST_IMAGES} test images")

# ==============================================================================
# Model Architecture: Pretrained EfficientNetB4
# ==============================================================================
with strategy.scope():
    pretrained_model = tf.keras.applications.EfficientNetB4(
        weights='imagenet',
        include_top=False,
        input_shape=[*IMAGE_SIZE, 3]
    )
    pretrained_model.trainable = True

    model = tf.keras.Sequential([
        pretrained_model,
        tf.keras.layers.GlobalAveragePooling2D(),
        tf.keras.layers.Dropout(0.3),
        tf.keras.layers.Dense(NUM_CLASSES, activation='softmax')
    ])

    model.compile(
        optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3),
        loss='sparse_categorical_crossentropy',
        metrics=['sparse_categorical_accuracy']
    )

print("\n--- Model Summary ---")
model.summary()

# Learning rate schedule
def lr_schedule(epoch):
    lr_start = 0.00001
    lr_max = 0.00005 * strategy.num_replicas_in_sync
    lr_min = 0.00001
    lr_ramp_ep = 4
    lr_sus_ep = 0
    lr_decay = 0.8
    if epoch < lr_ramp_ep:
        lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start
    elif epoch < lr_ramp_ep + lr_sus_ep:
        lr = lr_max
    else:
        lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min
    return lr

lr_callback = tf.keras.callbacks.LearningRateScheduler(lr_schedule, verbose=True)

# ==============================================================================
# Training Phase
# ==============================================================================
print(f"\n--- Training EfficientNetB4 for {EPOCHS} Epochs ---")
history = model.fit(
    get_training_dataset(),
    steps_per_epoch=STEPS_PER_EPOCH,
    epochs=EPOCHS,
    callbacks=[lr_callback]
)

# ==============================================================================
# Test Inference & Submission Generation
# ==============================================================================
print("\n--- Generating Predictions for Test Set ---")
test_ds = get_test_dataset(ordered=True)
test_images_ds = test_ds.map(lambda image, idnum: image)
probabilities = model.predict(test_images_ds)
predictions = np.argmax(probabilities, axis=-1)

test_ids_ds = test_ds.map(lambda image, idnum: idnum).unbatch()
test_ids = next(iter(test_ids_ds.batch(NUM_TEST_IMAGES))).numpy().astype('U')

sub_df = pd.DataFrame({'id': test_ids, 'label': predictions})
submission_file = "submission.csv"
sub_df.to_csv(submission_file, index=False)

print(f"\nSubmission successfully saved to {submission_file}! Shape: {sub_df.shape}")
print(sub_df.head(10))
