You need to agree to share your contact information to access this model

This repository is publicly accessible, but you have to accept the conditions to access its files and content.

Log in or Sign Up to review the conditions and access this model content.

Gemma-7B Survey Response Classifier - Inference

This document provides instructions for using the inference script to make predictions with the trained Gemma-7B model for survey response classification.

Requirements

  • Python 3.8+
  • PyTorch
  • Transformers
  • PEFT (Parameter-Efficient Fine-Tuning)
  • Pandas
  • BitsAndBytes (for 4-bit quantization)

Usage

The inference script (inference_gemma.py) supports two modes of operation:

  1. Single Prediction: Classify a single question-response pair
  2. Batch Prediction: Process multiple question-response pairs from a CSV file

Single Prediction

To classify a single question-response pair:

python inference_gemma.py --model_path /path/to/model --question "What is your favorite color?" --response "My favorite color is blue because it reminds me of the ocean."

Batch Prediction

To process multiple question-response pairs from a CSV file:

python inference_gemma.py --model_path /path/to/model --input_file data.csv --output_file predictions.csv --batch_size 4

The input CSV file should have at least the following columns:

  • question: The survey question
  • response: The response to classify

If the file also contains a label column, the script will calculate and report the accuracy of the predictions.

Command Line Arguments

  • --model_path: Path to the trained model directory (required)
  • --input_file: Path to CSV file with question and response columns (for batch prediction)
  • --output_file: Path to save predictions (default: 'predictions.csv')
  • --question: Question for single prediction
  • --response: Response for single prediction
  • --batch_size: Batch size for inference (default: 4)

Output Format

Single Prediction

The script will print the question, response, prediction (Positive/Negative), and the confidence probability.

Batch Prediction

The script will generate a CSV file with the following columns:

  • question: The original question
  • response: The original response
  • true_label: The true label (if available in the input file)
  • prediction: The predicted class (0 or 1)
  • probability: The confidence probability
  • label: The human-readable label (Positive/Negative)

Model Loading

The inference script is designed to handle different model formats:

  1. PEFT Adapter Models: If the model directory contains an adapter_config.json file, the script will load the model using the PEFT configuration.

  2. Direct Model Weights: The script will look for any .bin files or model.safetensors in the model directory and attempt to load them.

  3. Base Model: If no model weights are found, the script will use the base Gemma-7B model without fine-tuning.

Common Model File Formats

The script can load models saved in the following formats:

  • pytorch_model.bin: Standard PyTorch model file
  • model.safetensors: SafeTensors format (more secure)
  • Any other .bin file containing model weights

Troubleshooting Model Loading

If you see the warning "No model weights found at [path]. Using base model without fine-tuning", check the following:

  1. Correct Model Path: Make sure you're pointing to the correct directory where your trained model is saved.

  2. Model File Format: Ensure your model was saved with one of the supported formats. The script will log the files it finds in the model directory.

  3. Training Output: Check the output directory from your training script. The model should be saved in one of these locations:

    • best_model/: Directory containing the best model based on validation metrics
    • final_model/: Directory containing the final model after training
    • checkpoint_epoch_X/: Checkpoints saved after each epoch
  4. Model Saving: If you're using a custom training script, ensure it's saving the model weights correctly.

Notes

  • The model uses 4-bit quantization for memory efficiency
  • The script automatically handles loading the model and tokenizer
  • Logs are saved to a timestamped file for debugging purposes
  • The model path should point to the directory where the trained model is saved (e.g., 'best_model' or 'final_model' from the training script)
Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for halderavik/contextual-classifier

Base model

google/gemma-7b
Adapter
(9195)
this model