Spaces:
Sleeping
Sleeping
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import tensorflow as tf | |
| from sklearn.utils.class_weight import compute_class_weight | |
| from sklearn.metrics import classification_report, confusion_matrix | |
| import matplotlib.pyplot as plt | |
| import seaborn as sns | |
| BASE_DIR = Path("data") | |
| TRAIN_DIR = BASE_DIR / "train" | |
| VAL_DIR = BASE_DIR / "val" | |
| TEST_DIR = BASE_DIR / "test" | |
| IMAGE_SIZE = (224, 224) | |
| BATCH_SIZE = 32 | |
| SEED = 42 | |
| EPOCHS = 25 | |
| OUTPUT_DIR = Path("models") | |
| OUTPUT_DIR.mkdir(exist_ok=True) | |
| def load_datasets(): | |
| train_ds = tf.keras.utils.image_dataset_from_directory( | |
| TRAIN_DIR, | |
| labels="inferred", | |
| label_mode="int", | |
| image_size=IMAGE_SIZE, | |
| batch_size=BATCH_SIZE, | |
| shuffle=True, | |
| seed=SEED, | |
| ) | |
| val_ds = tf.keras.utils.image_dataset_from_directory( | |
| VAL_DIR, | |
| labels="inferred", | |
| label_mode="int", | |
| image_size=IMAGE_SIZE, | |
| batch_size=BATCH_SIZE, | |
| shuffle=False, | |
| ) | |
| test_ds = tf.keras.utils.image_dataset_from_directory( | |
| TEST_DIR, | |
| labels="inferred", | |
| label_mode="int", | |
| image_size=IMAGE_SIZE, | |
| batch_size=BATCH_SIZE, | |
| shuffle=False, | |
| ) | |
| class_names = train_ds.class_names | |
| autotune = tf.data.AUTOTUNE | |
| train_ds = train_ds.prefetch(autotune) | |
| val_ds = val_ds.prefetch(autotune) | |
| test_ds = test_ds.prefetch(autotune) | |
| return train_ds, val_ds, test_ds, class_names | |
| def compute_weights(): | |
| raw_train = tf.keras.utils.image_dataset_from_directory( | |
| TRAIN_DIR, | |
| labels="inferred", | |
| label_mode="int", | |
| image_size=IMAGE_SIZE, | |
| batch_size=BATCH_SIZE, | |
| shuffle=True, | |
| seed=SEED, | |
| ) | |
| labels = [] | |
| for _, y in raw_train.unbatch(): | |
| labels.append(int(y.numpy())) | |
| labels = np.array(labels) | |
| classes = np.arange(labels.max() + 1) | |
| weights = compute_class_weight( | |
| class_weight="balanced", | |
| classes=classes, | |
| y=labels, | |
| ) | |
| return {int(i): float(w) for i, w in enumerate(weights)} | |
| def build_model(num_classes): | |
| data_augmentation = tf.keras.Sequential( | |
| [ | |
| tf.keras.layers.RandomFlip("horizontal"), | |
| tf.keras.layers.RandomRotation(0.1), | |
| tf.keras.layers.RandomZoom(0.1), | |
| tf.keras.layers.RandomTranslation(0.05, 0.05), | |
| ], | |
| name="data_augmentation", | |
| ) | |
| rescale = tf.keras.layers.Rescaling(1.0 / 255.0, name="rescale") | |
| base_model = tf.keras.applications.MobileNetV2( | |
| input_shape=IMAGE_SIZE + (3,), | |
| include_top=False, | |
| weights="imagenet", | |
| ) | |
| base_model.trainable = False | |
| inputs = tf.keras.Input(shape=IMAGE_SIZE + (3,), name="image") | |
| x = data_augmentation(inputs) | |
| x = rescale(x) | |
| x = base_model(x, training=False) | |
| x = tf.keras.layers.GlobalAveragePooling2D()(x) | |
| x = tf.keras.layers.Dropout(0.3)(x) | |
| outputs = tf.keras.layers.Dense(num_classes, activation="softmax")(x) | |
| model = tf.keras.Model(inputs, outputs, name="agrovision_mobilenetv2") | |
| model.compile( | |
| optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4), | |
| loss=tf.keras.losses.SparseCategoricalCrossentropy(), | |
| metrics=[ | |
| "accuracy", | |
| tf.keras.metrics.SparseTopKCategoricalAccuracy( | |
| k=3, name="top3_accuracy" | |
| ), | |
| ], | |
| ) | |
| return model | |
| def plot_history(history, out_path): | |
| acc = history.history.get("accuracy", []) | |
| val_acc = history.history.get("val_accuracy", []) | |
| loss = history.history.get("loss", []) | |
| val_loss = history.history.get("val_loss", []) | |
| epochs_range = range(1, len(acc) + 1) | |
| plt.figure(figsize=(10, 4)) | |
| plt.subplot(1, 2, 1) | |
| plt.plot(epochs_range, acc, label="train_acc") | |
| plt.plot(epochs_range, val_acc, label="val_acc") | |
| plt.legend() | |
| plt.title("Accuracy") | |
| plt.subplot(1, 2, 2) | |
| plt.plot(epochs_range, loss, label="train_loss") | |
| plt.plot(epochs_range, val_loss, label="val_loss") | |
| plt.legend() | |
| plt.title("Loss") | |
| plt.tight_layout() | |
| plt.savefig(out_path) | |
| plt.close() | |
| def evaluate_on_test(model, test_ds, class_names, out_cm_path): | |
| y_true = [] | |
| y_pred = [] | |
| for x_batch, y_batch in test_ds: | |
| preds = model.predict(x_batch, verbose=0) | |
| y_true.extend(y_batch.numpy().tolist()) | |
| y_pred.extend(np.argmax(preds, axis=1).tolist()) | |
| y_true = np.array(y_true) | |
| y_pred = np.array(y_pred) | |
| print("\nClassification report:") | |
| print(classification_report(y_true, y_pred, target_names=class_names)) | |
| cm = confusion_matrix(y_true, y_pred) | |
| plt.figure(figsize=(14, 12)) | |
| sns.heatmap( | |
| cm, | |
| annot=False, | |
| cmap="Blues", | |
| xticklabels=class_names, | |
| yticklabels=class_names, | |
| ) | |
| plt.xlabel("Predicted") | |
| plt.ylabel("True") | |
| plt.xticks(rotation=90) | |
| plt.yticks(rotation=0) | |
| plt.tight_layout() | |
| plt.savefig(out_cm_path) | |
| plt.close() | |
| def main(): | |
| train_ds, val_ds, test_ds, class_names = load_datasets() | |
| num_classes = len(class_names) | |
| (OUTPUT_DIR / "class_names.json").write_text( | |
| json.dumps(class_names, indent=2), | |
| encoding="utf-8", | |
| ) | |
| class_weight = compute_weights() | |
| model = build_model(num_classes) | |
| checkpoint_path = OUTPUT_DIR / "agrovision_best.keras" | |
| callbacks = [ | |
| tf.keras.callbacks.ModelCheckpoint( | |
| filepath=str(checkpoint_path), | |
| monitor="val_accuracy", | |
| save_best_only=True, | |
| verbose=1, | |
| ), | |
| tf.keras.callbacks.EarlyStopping( | |
| monitor="val_accuracy", | |
| patience=5, | |
| restore_best_weights=True, | |
| verbose=1, | |
| ), | |
| tf.keras.callbacks.ReduceLROnPlateau( | |
| monitor="val_loss", | |
| factor=0.5, | |
| patience=3, | |
| verbose=1, | |
| ), | |
| ] | |
| history = model.fit( | |
| train_ds, | |
| validation_data=val_ds, | |
| epochs=EPOCHS, | |
| class_weight=class_weight, | |
| callbacks=callbacks, | |
| ) | |
| hist_path = OUTPUT_DIR / "training_curves.png" | |
| plot_history(history, hist_path) | |
| print("\nTest evaluation:") | |
| test_metrics = model.evaluate(test_ds, verbose=1) | |
| for name, value in zip(model.metrics_names, test_metrics): | |
| print(f"{name}: {value:.4f}") | |
| cm_path = OUTPUT_DIR / "confusion_matrix.png" | |
| evaluate_on_test(model, test_ds, class_names, cm_path) | |
| final_path = OUTPUT_DIR / "agrovision_final.keras" | |
| model.save(final_path) | |
| print(f"\nSaved final model to: {final_path}") | |
| if __name__ == "__main__": | |
| main() | |