Код: Выделить всё
local_dataset_path = 'cifar-10-python.tar.gz'
extracted_path = 'cifar-10-batches-py'
# Extract the .tar.gz file
with tarfile.open(local_dataset_path, 'r:gz') as file:
file.extractall() # Extracts into the current directory
# Function to load data from a batch file
def load_batch(batch_path):
with open(batch_path, 'rb') as file:
batch = pickle.load(file, encoding='latin1')
images = batch['data']
labels = batch['labels']
# Reshape and normalize the images
images = images.reshape((len(images), 3, 32, 32)).transpose(0, 2, 3, 1) / 255.0
labels = np.array(labels)
return images, labels
# Load all training batches
training_images = []
training_labels = []
for i in range(1, 6):
batch_path = os.path.join(extracted_path, f'data_batch_{i}')
images, labels = load_batch(batch_path)
training_images.append(images)
training_labels.append(labels)
# Concatenate all training images and labels
training_images = np.concatenate(training_images)
training_labels = np.concatenate(training_labels)
# Load the test batch
test_batch_path = os.path.join(extracted_path, 'test_batch')
testing_images, testing_labels = load_batch(test_batch_path)
# Class names for the CIFAR-10 dataset
class_names = ['Plane', 'Car', 'Bird', 'Cat', 'Deer', 'Dog', 'Frog', 'Horse', 'Ship', 'Truck']
# Plot the first 16 images
import matplotlib.pyplot as plt
for image in range(16):
plt.subplot(4, 4, image + 1)
plt.xticks([])
plt.yticks([])
plt.imshow(training_images[image])
plt.xlabel(class_names[training_labels[image]])
# plt.show()
training_images = training_images[:2000]
training_labels = training_labels[:2000]
testing_images = testing_images[:4000]
testing_labels = testing_labels[:4000]
model = models.Sequential()
model.add(layers.Conv2D(32, (3, 3), activation='relu', input_shape=(32, 32, 3)))
model.add(layers.MaxPooling2D((2,2)))
model.add(layers.Conv2D(64, (3,3), activation='relu'))
model.add(layers.MaxPooling2D((2,2)))
model.add(layers.Conv2D(64, (3,3), activation='relu'))
model.add(layers.Flatten())
model.add(layers.Dense(64, activation='relu'))
model.add(layers.Dense(10, activation = 'softmax'))
model.compile(loss='sparse_categorical_crossentropy', optimizer='adam', metrics = ['accuracy'])
model.fit(training_images, training_labels, epochs=10, validation_data=(testing_images, testing_labels))
loss, accuracy = model.evaluate(testing_images, testing_labels)
print(f"Loss: {loss}")
print(f"Accuracy: {accuracy}")
model.save('image_classifier.model')
Код: Выделить всё
ValueError: Invalid filepath extension for saving. Please add either a `.keras` extension
for the native Keras format (recommended) or a `.h5` extension. Use `model.export(filepath)`
if you want to export a SavedModel for use with TFLite/TFServing/etc.
Received: filepath=image_classifier.model.
Подробнее здесь: https://stackoverflow.com/questions/790 ... -cnn-model