CoolFace
Modelpublic

hula07/cifar

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes6downloads
Hugging.py174 linesDownload Raw Back to root
1import numpy as np2import tensorflow as tf3from PIL import Image4import torch5from datasets import load_metric6from datasets import load_dataset7from transformers import (ViTFeatureExtractor, ViTForImageClassification, TrainingArguments, Trainer, create_optimizer)8 9def convert_to_tf_tensor(image: Image):10    # np_image = np.array(image)11    # tf_image = tf.convert_to_tensor(np_image)12    # return tf.expand_dims(tf_image, 0)13    np_image = np.array(image)14    tf_image = tf.convert_to_tensor(np_image)15    tf_image = tf.image.resize(tf_image, [224, 224])  # Resize to 224x22416    tf_image = tf.repeat(tf_image, 3, -1)  # Repeat along the color dimension to simulate 3 channels17    return tf.expand_dims(tf_image, 0)18 19def preprocess(batch):20    # take a list of PIL images and turn them to pixel values21    inputs = feature_extractor(22        batch['img'],23        return_tensors='pt'24    )25    # include the labels26    inputs['label'] = batch['label']27    return inputs28 29def collate_fn(batch):30    return {31        'pixel_values': torch.stack([x['pixel_values'] for x in batch]),32        'labels': torch.tensor([x['label'] for x in batch])33    }34 35def compute_metrics(p):36        return metric.compute(37        predictions=np.argmax(p.predictions, axis=1),38        references=p.label_ids39    )40 41if __name__ == '__main__':42    dataset_train = load_dataset(43        'cifar10',44        split='train[:1000]',  # training dataset45        ignore_verifications=False  # set to True if seeing splits Error46    )47    print(dataset_train)48 49    dataset_test = load_dataset(50        'cifar10',51        split='test',  # training dataset52        ignore_verifications=True  # set to True if seeing splits Error53    )54    print(dataset_test)55 56    # check how many labels/number of classes57    num_classes = len(set(dataset_train['label']))58    labels = dataset_train.features['label']59    print(num_classes, labels)60 61    print(dataset_train[0]['label'], labels.names[dataset_train[0]['label']])62    # import model63    model_id = 'google/vit-base-patch16-224-in21k'64    feature_extractor = ViTFeatureExtractor.from_pretrained(65        model_id66    )67    print(feature_extractor)68 69    example = feature_extractor(70        dataset_train[0]['img'],71        return_tensors='pt'72    )73    print(example)74    print(example['pixel_values'].shape)75 76    # transform the training dataset77    prepared_train = dataset_train.with_transform(preprocess)78    prepared_test = dataset_test.with_transform(preprocess)79 80    # accuracy metric81    metric = load_metric("accuracy")82 83    training_args = TrainingArguments(84        output_dir="./cifar",85        per_device_train_batch_size=16,86        evaluation_strategy="steps",87        num_train_epochs=4,88        save_steps=100,89        eval_steps=100,90        logging_steps=10,91        learning_rate=2e-4,92        save_total_limit=2,93        remove_unused_columns=False,94        push_to_hub=True,95        load_best_model_at_end=True,96        # output_dir='./cifar',97        # per_device_train_batch_size=16,98        # evaluation_strategy='steps',99        # num_train_epochs=4,100        # save_steps=100,101        # eval_steps=100,102        # logging_steps=10,103        # learning_rate=2e-4,104        # save_total_limit=2,105        # remove_unused_columns=False,106        # push_to_hub=True,107        # push_to_hub_model_id="classify_images",108        # load_best_model_at_end=True,109    )110 111    labels = dataset_train.features['label'].names112 113    model = ViTForImageClassification.from_pretrained(114        model_id,  # classification head115        num_labels=len(labels)116    )117 118    trainer = Trainer(119        model=model,120        args=training_args,121        data_collator=collate_fn,122        compute_metrics=compute_metrics,123        train_dataset=prepared_train,124        eval_dataset=prepared_test,125        tokenizer=feature_extractor,126    )127 128    # Run the training129    train_results = trainer.train()130    trainer.push_to_hub()131    # save tokenizer with the model132    trainer.save_model()133    trainer.log_metrics("train", train_results.metrics)134    trainer.save_metrics("train", train_results.metrics)135    # save the trainer state136    trainer.save_state()137    batch_size = 16138    num_epochs = 5139    num_train_steps = len(dataset_train["train"]) * num_epochs140    learning_rate = 3e-5141    weight_decay_rate = 0.01142 143    optimizer, lr_schedule = create_optimizer(144        init_lr=learning_rate,145        num_train_steps=num_train_steps,146        weight_decay_rate=weight_decay_rate,147        num_warmup_steps=0,148    )149    tf_train_dataset = prepared_train.to_tf_dataset(150        features=["pixel_values"],151        labels=["label"],152        batch_size=batch_size,153        shuffle=True,154        collate_fn=collate_fn155    )156 157    tf_eval_dataset = prepared_test.to_tf_dataset(158        features=["pixel_values"],159        labels=["label"],160        batch_size=batch_size,161        shuffle=False,162        collate_fn=collate_fn163    )164    loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)165    model.compile(optimizer=optimizer, loss=loss)166 167    metrics = trainer.evaluate(prepared_test)168    trainer.log_metrics("eval", metrics)169    trainer.save_metrics("eval", metrics)170    # Evaluate the model171    eval_results = trainer.evaluate()172 173    print(eval_results)174