【问题标题】:MNIST Dataset StructureMNIST 数据集结构
【发布时间】:2019-05-14 05:53:04
【问题描述】:

我下载了一个实现遗传算法的代码。它使用默认数据集mnist。我想更改默认数据集“mnist”,但同时我想知道数据集的结构,以便我可以按照mnist 的方式格式化我的数据。例如,我还想知道,如果我已经格式化了我的数据my_own_data_set,那么输入函数调用network.train(my_own_data_set) 是否有效?函数train.network接受什么数据类型?

"""Entry point to evolving the neural network. Start here."""
import logging
from optimizer import Optimizer
from tqdm import tqdm

# Setup logging.
logging.basicConfig(
    format='%(asctime)s - %(levelname)s - %(message)s',
    datefmt='%m/%d/%Y %I:%M:%S %p',
    level=logging.DEBUG,
    filename='log.txt'
)

def train_networks(networks, dataset):
    """Train each network.

    Args:
        networks (list): Current population of networks
        dataset (str): Dataset to use for training/evaluating
   """
   pbar = tqdm(total=len(networks))
   for network in networks:
       network.train(dataset)
       pbar.update(1)
       pbar.close()

   def get_average_accuracy(networks):
       """Get the average accuracy for a group of networks.

   Args:
       networks (list): List of networks

   Returns:
       float: The average accuracy of a population of networks.

   """
   total_accuracy = 0
   for network in networks:
       total_accuracy += network.accuracy

   return total_accuracy / len(networks)

   def generate(generations, population, nn_param_choices, dataset):
       """Generate a network with the genetic algorithm.

   Args:
       generations (int): Number of times to evole the population
       population (int): Number of networks in each generation
       nn_param_choices (dict): Parameter choices for networks
       dataset (str): Dataset to use for training/evaluating

   """
   optimizer = Optimizer(nn_param_choices)
   networks = optimizer.create_population(population)

   # Evolve the generation.
   for i in range(generations):
       logging.info("***Doing generation %d of %d***" %
                 (i + 1, generations))

       # Train and get accuracy for networks.
       train_networks(networks, dataset)

       # Get the average accuracy for this generation.
       average_accuracy = get_average_accuracy(networks)

       # Print out the average accuracy each generation.
       logging.info("Generation average: %.2f%%" % (average_accuracy * 100))
       logging.info('-'*80)

       # Evolve, except on the last iteration.
       if i != generations - 1:
           # Do the evolution.
           networks = optimizer.evolve(networks)

    # Sort our final population.
    networks = sorted(networks, key=lambda x: x.accuracy, reverse=True)

    # Print out the top 5 networks.
    print_networks(networks[:5])

def print_networks(networks):
    """Print a list of networks.

    Args:
         networks (list): The population of networks

    """
    logging.info('-'*80)
    for network in networks:
        network.print_network()

def main():
    """Evolve a network."""
    generations = 10  # Number of times to evole the population.
    population = 20  # Number of networks in each generation.
    dataset = 'mnist'

    nn_param_choices = {
        'nb_neurons': [64, 128, 256, 512, 768, 1024],
        'nb_layers': [1, 2, 3, 4],
        'activation': ['relu', 'elu', 'tanh', 'sigmoid'],
        'optimizer': ['rmsprop', 'adam', 'sgd', 'adagrad',
                  'adadelta', 'adamax', 'nadam'],
    }

    logging.info("***Evolving %d generations with population %d***" %
             (generations, population))

    generate(generations, population, nn_param_choices, dataset)

if __name__ == '__main__':
     main()

【问题讨论】:

  • “格式”是什么意思?
  • @cheersmate -> 我的意思是我想预处理我的数据,MNIST 在数据集中的结构方式。我想知道 MNIST 数据集的维度、它拥有的数据类型等等。
  • 所以您可以在使用 MNIST 运行时首先检查数组形状、dtypes 等,然后为您自己的数据集添加一些检查?
  • @cheersmate 但数据集变量只是一个字符串,请参见上面的代码。数据集 = 'mnist'。我不知道 network.train(dataset) 如何提取数据进行训练。所以我无法检查它的尺寸。
  • 你有源代码来了解内部发生了什么吗?

标签: python machine-learning neural-network genetic-algorithm mnist


【解决方案1】:

每一行都是一个字符串数组,785 个字符串用引号括起来,用逗号分隔。 训练集中大约 60,000 行,测试集中大约 10,000 行。

该行以图片中该行其余部分的标签开始...

“8”、“0”、“0”、“0”、“0”、“0”、“0”、“0”、“0”、“0”、“0”、“0” , ... "255"...

所以第一个字符串项是图像代表的数字,其余是代表图像的 0 - 255 的 784 个字符串,这是一个 28x28 的图像,以行首尾相连的长字符串数组表示。

string[785] 行,每行

【讨论】:

    猜你喜欢
    • 2018-05-13
    • 1970-01-01
    • 2021-04-08
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2014-12-02
    • 2021-03-04
    • 2016-06-08
    相关资源
    最近更新 更多