import cv2
import numpy as np
from tensorflow.keras.datasets import mnist
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense, Flatten

# ====== 1. TRENING MODELU ======
(x_train, y_train), (x_test, y_test) = mnist.load_data()

x_train = x_train / 255.0
x_test = x_test / 255.0

model = Sequential([
    Flatten(input_shape=(28, 28)),
    Dense(128, activation="relu"),
    Dense(10, activation="softmax")
])

model.compile(
    optimizer="adam",
    loss="sparse_categorical_crossentropy",
    metrics=["accuracy"]
)

model.fit(x_train, y_train, epochs=3, verbose=1)

# ====== 2. OKNO DO RYSOWANIA ======
canvas = np.zeros((280, 280), dtype=np.uint8)

drawing = False

def draw(event, x, y, flags, param):
    global drawing

    if event == cv2.EVENT_LBUTTONDOWN:
        drawing = True

    elif event == cv2.EVENT_MOUSEMOVE:
        if drawing:
            cv2.circle(canvas, (x, y), 10, 255, -1)

    elif event == cv2.EVENT_LBUTTONUP:
        drawing = False

cv2.namedWindow("Rysuj cyfrę")
cv2.setMouseCallback("Rysuj cyfrę", draw)

print("Rysuj cyfrę myszką. Naciśnij 'p' aby przewidzieć, 'c' aby wyczyścić, 'q' aby wyjść.")

# ====== 3. PĘTLA ======
while True:
    cv2.imshow("Rysuj cyfrę", canvas)

    key = cv2.waitKey(1)

    if key == ord('c'):
        canvas[:] = 0

    elif key == ord('p'):
        # przygotowanie obrazu
        img = cv2.resize(canvas, (28, 28))
        img = img / 255.0
        img = img.reshape(1, 28, 28)

        prediction = model.predict(img)
        print("👉 AI zgaduje:", np.argmax(prediction))

    elif key == ord('q'):
        break

cv2.destroyAllWindows()