krisha06's picture
Update app.py
5a3c5c6 verified
Raw
History Blame Contribute Delete
2.36 kB
import streamlit as st
from ultralytics import YOLO
import os
import zipfile
# Function to unzip the dataset (if not already extracted)
def unzip_dataset(zip_path, extract_path):
if not os.path.exists(extract_path):
os.makedirs(extract_path)
with zipfile.ZipFile(zip_path, 'r') as zip_ref:
zip_ref.extractall(extract_path)
st.success("Dataset extracted successfully!")
# Paths to the dataset and model weights
zip_path = 'Clothing Segmentation DataSet.v9i.yolov8.zip' # File should be uploaded in root
extract_path = '/tmp/clothing_dataset' # /tmp is writable
model_weights = '/tmp/clothing_model.pt' # Place to save/load model weights if needed
# Check if the model is already trained or needs to be trained
model_trained = os.path.exists(model_weights)
if not model_trained:
# Unzip dataset if necessary
unzip_dataset(zip_path, extract_path)
# Train the model
model = YOLO('yolov8s.pt') # Using YOLOv8 small model for faster training
# Path to the data.yaml
data_yaml_path = os.path.join(extract_path, 'data.yaml')
# Train the model
with st.spinner('Training model...'):
model.train(data=data_yaml_path, epochs=50, imgsz=640, batch=16)
model.save(model_weights) # Save the model after training
st.success("Model trained and saved successfully!")
# Load the trained YOLOv8 model
model = YOLO(model_weights).to('cpu')
# Streamlit App Interface
st.title("Clothing Item Segmentation")
# Upload an image for segmentation
uploaded_image = st.file_uploader("Upload an image of clothing...", type=["jpg", "png", "jpeg"])
if uploaded_image is not None:
# Display uploaded image
st.image(uploaded_image, caption="Uploaded Image", use_column_width=True)
# Run inference using the trained YOLOv8 model
st.spinner("Running segmentation...")
results = model.predict(source=uploaded_image) # Run inference
# Show the segmented image result
st.image(results[0].plot(), caption="Segmentation Result", use_column_width=True)
# Display the prediction details (class names and confidence)
st.write("Prediction Details:")
for result in results[0].boxes.data:
class_id = int(result[5]) # Class ID
confidence = result[4].item() # Confidence
st.write(f"Class ID: {class_id}, Confidence: {confidence:.2f}")