hula07/cifar
06
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 