from __future__ import annotations
import argparse
from pathlib import Path
import tensorflow as tf


def parse_args():
    parser = argparse.ArgumentParser(description="TensorFlow CIFAR-10 CNN with checkpointing")
    parser.add_argument('--batch-size', type=int, default=64, metavar='N',
                        help='input batch size for training (default: 64)')
    parser.add_argument('--test-batch-size', type=int, default=1000, metavar='N',
                        help='input batch size for testing (default: 1000)')
    parser.add_argument('--epochs', type=int, default=10, metavar='N',
                        help='number of epochs to train (default: 10)')
    parser.add_argument('--lr', type=float, default=0.01, metavar='LR',
                        help='learning rate (default: 0.01)')
    parser.add_argument('--momentum', type=float, default=0.5, metavar='M',
                        help='SGD momentum (default: 0.5)')
    parser.add_argument('--no-cuda', action='store_true', default=False,
                        help='disables CUDA training')
    parser.add_argument('--seed', type=int, default=1, metavar='S',
                        help='random seed (default: 1)')
    parser.add_argument('--log-interval', type=int, default=10, metavar='N',
                        help='how many batches to wait before logging training status')
    parser.add_argument('--ckpt-dir', required=True, help='path to save and load checkpoints')
    parser.add_argument('--resume-training', action='store_true', help='resume training from latest checkpoint')
    return parser.parse_args()


def disable_gpu_if_requested(no_cuda: bool) -> None:
    if not no_cuda:
        return
    try:
        tf.config.set_visible_devices([], 'GPU')
        print('GPU devices disabled via --no-cuda; using CPU.')
    except Exception as exc:  # pragma: no cover - defensive
        print(f'Could not disable GPU devices: {exc}')


def build_datasets(batch_size: int, test_batch_size: int, seed: int):
    (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.reshape(-1)
    y_test = y_test.reshape(-1)

    train_ds = (tf.data.Dataset.from_tensor_slices((x_train, y_train))
                .shuffle(buffer_size=10000, seed=seed)
                .batch(batch_size)
                .prefetch(tf.data.AUTOTUNE))

    test_ds = (tf.data.Dataset.from_tensor_slices((x_test, y_test))
               .batch(test_batch_size)
               .prefetch(tf.data.AUTOTUNE))
    return train_ds, test_ds


def create_model():
    return tf.keras.Sequential([
        tf.keras.layers.Input(shape=(32, 32, 3)),
        tf.keras.layers.Conv2D(32, 3, activation='relu'),
        tf.keras.layers.MaxPooling2D(pool_size=(2, 2)),
        tf.keras.layers.Conv2D(64, 3, activation='relu'),
        tf.keras.layers.Dropout(0.3),
        tf.keras.layers.MaxPooling2D(pool_size=(2, 2)),
        tf.keras.layers.Flatten(),
        tf.keras.layers.Dense(50, activation='relu'),
        tf.keras.layers.Dropout(0.5),
        tf.keras.layers.Dense(10)
    ])


def main():
    args = parse_args()
    tf.random.set_seed(args.seed)
    disable_gpu_if_requested(args.no_cuda)

    ckpt_dir = Path(args.ckpt_dir)
    ckpt_dir.mkdir(parents=True, exist_ok=True)

    train_ds, test_ds = build_datasets(args.batch_size, args.test_batch_size, args.seed)

    model = create_model()
    loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
    optimizer = tf.keras.optimizers.SGD(learning_rate=args.lr, momentum=args.momentum)

    ckpt = tf.train.Checkpoint(epoch=tf.Variable(0, dtype=tf.int64),
                               optimizer=optimizer,
                               model=model)
    manager = tf.train.CheckpointManager(ckpt, ckpt_dir, max_to_keep=5)

    start_epoch = 0
    if args.resume_training and manager.latest_checkpoint:
        ckpt.restore(manager.latest_checkpoint)
        start_epoch = int(ckpt.epoch.numpy())
        print(f'Resuming from checkpoint: {manager.latest_checkpoint} (start at epoch {start_epoch + 1})')
    elif args.resume_training:
        print(f'No checkpoints found in {ckpt_dir}. Starting from scratch...')

    @tf.function
    def train_step(images, labels):
        with tf.GradientTape() as tape:
            logits = model(images, training=True)
            loss_value = loss_fn(labels, logits)
        grads = tape.gradient(loss_value, model.trainable_variables)
        optimizer.apply_gradients(zip(grads, model.trainable_variables))
        return loss_value, logits

    @tf.function
    def test_step(images, labels):
        logits = model(images, training=False)
        t_loss = loss_fn(labels, logits)
        return t_loss, logits

    for epoch in range(start_epoch, args.epochs):
        train_loss = tf.keras.metrics.Mean()
        train_accuracy = tf.keras.metrics.SparseCategoricalAccuracy()
        test_loss = tf.keras.metrics.Mean()
        test_accuracy = tf.keras.metrics.SparseCategoricalAccuracy()

        for batch_idx, (images, labels) in enumerate(train_ds):
            loss_value, logits = train_step(images, labels)
            train_loss.update_state(loss_value)
            train_accuracy.update_state(labels, logits)
            if batch_idx % args.log_interval == 0:
                print(f'Train Epoch: {epoch + 1} [{batch_idx * len(images):>5d}]\tLoss: {loss_value.numpy():.6f}')

        for images, labels in test_ds:
            t_loss, logits = test_step(images, labels)
            test_loss.update_state(t_loss)
            test_accuracy.update_state(labels, logits)

        ckpt.epoch.assign(epoch + 1)
        ckpt_path = manager.save()
        print(f'Saved checkpoint to {ckpt_path}')

        print(f'\nEpoch {epoch + 1}/{args.epochs}: ')
        print(f'  Train   - loss: {train_loss.result():.4f}, accuracy: {train_accuracy.result() * 100:.2f}%')
        print(f'  Test    - loss: {test_loss.result():.4f}, accuracy: {test_accuracy.result() * 100:.2f}%\n')


if __name__ == '__main__':
    main()
