关于Keras文档中MNIST示例神经网络预测逻辑的技术问询
Great job getting through the neural network crash course and diving into MNIST with Keras in R—let’s break down your questions one by one to get you fully up to speed:
1. What's the difference between x_train and y_train?
Think of this as the "input" vs "answer key" for your model:
x_trainholds the feature data: each entry is a 28x28 matrix of pixel values (0-255, normalized to 0-1 in most setups) representing a handwritten digit. This is what the model looks at to learn patterns.y_trainholds the label data: each entry is the actual digit (0-9) that corresponds to the image inx_train. This is the ground truth the model uses to adjust its weights and learn how to map pixel patterns to the correct number.
For example, x_train[1,,] is the pixel matrix of the first training image, and y_train[1] is the real digit that image represents.
2. Why does x_test look like all zeros?
This is almost certainly due to data normalization! When loading MNIST with Keras, the pixel values are automatically scaled from their original 0-255 range (raw grayscale) down to 0-1. Most of the image is background (white, originally 255), so those values become 1? Wait no—wait, actually, MNIST uses white digits on a black background, so normalized background values are 0, and digit pixels are closer to 1.
To confirm, run range(x_test)—you’ll see values between 0 and 1. To view the original raw pixel values, just multiply by 255: x_test_raw <- x_test * 255, and you’ll see the non-zero values for the digit pixels.
3. How to view the original image to verify predictions?
You can use R’s base plotting tools or ggplot2 to visualize the image. Here are two simple methods:
Using base R:
# Grab the first test sample and reshape it to 28x28 img_matrix <- matrix(x_test[1,,], nrow = 28, ncol = 28) # Plot the image (reverse y-axis to fix the orientation) image(img_matrix, col = gray.colors(255), axes = FALSE, ylim = c(1, 0))
Using ggplot2 (for a cleaner look):
library(ggplot2) library(tidyr) # Convert the pixel matrix to a data frame for ggplot img_df <- as.data.frame(img_matrix) %>% mutate(y = 1:28) %>% pivot_longer(-y, names_to = "x", values_to = "intensity") %>% mutate(x = as.integer(gsub("V", "", x))) # Plot as a heatmap ggplot(img_df, aes(x = x, y = y, fill = intensity)) + geom_tile() + scale_fill_gradient(low = "white", high = "black") + scale_y_reverse() + theme_void()
4. How to draw your own digits and predict them?
Here’s a step-by-step workflow to create, preprocess, and predict custom handwritten digits:
Step 1: Draw and save a digit
Create a 28x28 blank image and draw your digit with base R’s plotting tools:
# Create a 28x28 PNG file with a white background png("my_custom_digit.png", width = 28, height = 28, bg = "white") par(mar = c(0, 0, 0, 0)) # Remove margins plot(0:28, 0:28, type = "n", axes = FALSE) # Blank plot # Draw your digit (example: drawing a 7) lines(c(5, 23), c(23, 23), lwd = 3) # Top horizontal line lines(c(5, 5), c(23, 5), lwd = 3) # Left vertical line lines(c(5, 23), c(5, 5), lwd = 3) # Bottom horizontal line dev.off() # Save and close the image
Step 2: Preprocess the image for your model
You need to match the format of the MNIST training data:
library(png) # Read the image img <- readPNG("my_custom_digit.png") # Convert to grayscale if it's a color image if (dim(img)[3] == 3) { img <- img[,,1] } # Invert colors (MNIST uses black digits on white background; we drew white on black) img <- 1 - img # Reshape to match the model's input shape (1 sample, 28x28, 1 channel) img_array <- array(img, dim = c(1, 28, 28, 1)) # Normalize to 0-1 (same as training data) img_array <- img_array / 255
Step 3: Predict with your trained model
# Get predicted class predicted_digit <- predict_classes(model, img_array) # Or get probability scores and pick the highest one probabilities <- predict(model, img_array) predicted_digit <- which.max(probabilities) - 1 # Adjust for 0-based indexing cat("Your custom digit is predicted as:", predicted_digit, "\n")
内容的提问来源于stack exchange,提问作者user3078100

