{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"},{"sourceId":1205039,"sourceType":"datasetVersion","datasetId":685665}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install efficientnet_pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-19T17:50:57.552083Z","iopub.execute_input":"2025-03-19T17:50:57.552417Z","iopub.status.idle":"2025-03-19T17:51:00.855133Z","shell.execute_reply.started":"2025-03-19T17:50:57.552388Z","shell.execute_reply":"2025-03-19T17:51:00.85405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport numpy as np\nimport cv2\nfrom efficientnet_pytorch import EfficientNet\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\n# Set random seed for reproducibility\nSEED = 42\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = True\n\n# Define transformations for validation/inference\ndef get_valid_transforms():\n    return A.Compose([\n        A.Resize(height=512, width=512, p=1.0),\n        ToTensorV2(p=1.0),\n    ], p=1.0)\n\n# Define a dataset class for single image inference\nclass SingleImageDataset(Dataset):\n    def __init__(self, image_path, transforms=None):\n        super().__init__()\n        self.image_path = image_path\n        self.transforms = transforms\n\n    def __getitem__(self, index: int):\n        image = cv2.imread(self.image_path, cv2.IMREAD_COLOR)\n        if image is None:\n            raise ValueError(f\"Failed to read image at {self.image_path}. Check if the file is a valid image.\")\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {'image': image}\n            sample = self.transforms(**sample)\n            image = sample['image']\n        return image\n\n    def __len__(self) -> int:\n        return 1\n\n# Define a function to classify an image\ndef classify_image(image_path, model, device):\n    # Check if the image file exists\n    if not os.path.exists(image_path):\n        raise FileNotFoundError(f\"Image file not found at {image_path}. Please provide the correct path.\")\n\n    # Load the image\n    dataset = SingleImageDataset(image_path, transforms=get_valid_transforms())\n    data_loader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=2, drop_last=False)\n\n    # Perform inference\n    model.eval()\n    with torch.no_grad():\n        for images in data_loader:\n            images = images.to(device).float()\n            outputs = model(images)\n            probs = nn.functional.softmax(outputs, dim=1).cpu().numpy()[0]\n    \n    # Determine the class and subclass\n    class_names = ['Cover', 'JMiPOD', 'JUNIWARD', 'UERD']\n    predicted_class_idx = np.argmax(probs)\n    predicted_class = class_names[predicted_class_idx]\n    is_stego = predicted_class != 'Cover'\n    \n    return is_stego, predicted_class, probs\n\n# Load the pre-trained model\ndef get_net():\n    net = EfficientNet.from_pretrained('efficientnet-b2')\n    net._fc = nn.Linear(in_features=1408, out_features=4, bias=True)\n    return net\n\n# Load the model checkpoint\ncheckpoint_path = '/kaggle/input/alaska2-public-baseline/best-checkpoint-023epoch.bin'  # Replace with your checkpoint path\nnet = get_net()\n\n# Check if the checkpoint file exists\nif not os.path.exists(checkpoint_path):\n    raise FileNotFoundError(f\"Checkpoint file not found at {checkpoint_path}. Please provide the correct path.\")\n\n# Load the checkpoint\ncheckpoint = torch.load(checkpoint_path, map_location=device)\nnet.load_state_dict(checkpoint['model_state_dict'])\nnet.eval()\n\n# Set the device (GPU if available, otherwise CPU)\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nnet.to(device)\n\n# Loop for user input\nwhile True:\n    # Ask the user for the image path\n    image_path = input(\"Enter the path to the image (or type 'exit' to quit): \").strip()\n\n    # Exit the loop if the user types 'exit'\n    if image_path.lower() == 'exit':\n        print(\"Exiting the program. Goodbye!\")\n        break\n\n    try:\n        # Classify the image\n        is_stego, predicted_class, probs = classify_image(image_path, net, device)\n\n        # Print the results\n        print(\"\\nClassification Results:\")\n        print(f\"Is Stego: {is_stego}\")\n        print(f\"Predicted Class: {predicted_class}\")\n        print(f\"Class Probabilities: {probs}\\n\")\n    except Exception as e:\n        # Handle errors (e.g., invalid file path or corrupted image)\n        print(f\"Error: {e}\\n\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-19T17:51:00.85677Z","iopub.execute_input":"2025-03-19T17:51:00.857102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}