torchnet:实践指南

作者:袖梨 2026-09-11

实际评估torchnet时,我先确认它解决的具体问题:基于 Torch 的机器学习工具库,提供数据集、训练引擎与评估组件。团队若要把它用于部署与运行环境,应先处理权限、依赖和环境差异会放大维护成本,否则试用结果很容易失真。短测时我会在非生产环境复现一次安装与运行,并保留依赖锁定、权限边界、日志、回滚和资源消耗的结果,方便团队复盘。我的判断是,它更适合愿意维护环境并重视故障恢复的工程团队;若眼下没有这类需求,先保留观察即可。

火炬网

torchnet 是 torch 的框架,它提供了一组 旨在鼓励代码重用以及鼓励 模块化编程。

目前,torchnet 提供了四组重要的类:

  • Dataset:以各种方式处理和预处理数据。
  • Engine:training/testing 机器学习算法。
  • Meter:仪表性能或任何其他数量。
  • Log:以一致的方式将性能或任何其他字符串输出到文件/磁盘。

有关 torchnet 框架的概述,另请参阅 本文。

安装

请先安装 torch,按照以下说明进行操作 torch.ch。 如果火炬是 已经安装,请确保您拥有最新版本 argcheck,否则你会得到 运行时出现奇怪的错误。

假设 torch 已经安装,torchnet 核心只是一组 lua 文件,因此使用 luarocks 安装它很简单

luarocks install torchnet

要运行本文中的 MNIST 示例,请安装 mnist 包:

luarocks install mnist

cd 进入已安装的 torchnet 包目录并运行:

th example/mnist.lua

文档

要求 torchnet 返回一个包含所有 torchnet 的局部变量 类构造函数。

local tnt = require 'torchnet'

tnt.Dataset()

torchnet提供了多种数据容器,可以方便地 相互之间插入,允许用户轻松连接、拆分、 批处理、重新采样等...数据集。

tnt.Dataset() 的实例 dataset 实现了两个主要方法:

  • dataset:size() 返回数据集的大小。
  • dataset:get(idx),其中 idx 是 1 到数据集大小之间的数字。

虽然使用 for 循环迭代数据集很容易,但有几个 尽管如此,还是提供了 DatasetIterator 迭代器,允许用户 以即时方式过滤掉一些样本,或者轻松并行化 数据获取。

torchnet中,dataset:get()返回的样本应该是Lua table。表的字段可以是任意的,即使有许多数据集 仅适用于火炬张量。

tnt.utils

Torchnet 提供了一组在 torchnet 上使用的 util 函数。

tnt.utils.table.clone(表)

该函数对表进行深层复制。

tnt.utils.table.merge(dst, src)

({
   dst = table  --
   src = table  --
})

该函数添加到目标表dest, 源表 source 中包含的元素。

副本很浅。

如果两个表中都存在某个键,则源表中的元素 是优选的。

tnt.utils.table.foreach(tbl, 闭包[, 递归])

({
   tbl       = table     --
   closure   = function  --
  [recursive = boolean]  --  [default=false]
})

该函数将closure定义的函数应用于 表 tbl

如果给出 recursive 并设置为 true,则 closure 函数 将递归地应用于表。

tnt.utils.table.canmergetensor(表)

检查表是否可以合并为张量。

tnt.utils.table.mergetensor(表)

({
   tbl = table  --
})

将表合并为一个额外维度的张量。

tnt.transform

Torchnet 提供了一组通用数据转换。 这些转换要么直接在数据上进行(e.g.,标准化) 或关于它们的结构。这个特别方便 操作 tnt.Dataset 时。

大多数转换都很简单,但可以由 组成 或 合并了。

transform.identity(...)

恒等变换接受任何输入并按原样返回。

例如,这个函数在编写时很有用 对来自多个来源和某些来源的数据进行转换 不得改造。

transform.compose(变换)

({
   transforms = table  --
})

该函数采用 table 函数, 组合它们以返回一个转换。

该函数假设转换表 由从 1 开始的连续有序键进行索引。 变换按升序排列。

例如,以下代码:

> f = transform.compose{
        [1] = function(x) return 2*x end,
        [2] = function(x) return x + 10 end,
        foo = function(x) return x / 2 end,
        [4] = function(x) return x - x end
   }
   > f(3)
   16

相当于组合[1]和[2]中存储的变换,i.e。, 定义以下变换:

> f =  function(x) return 2*x + 10 end

请注意,使用键 foo4 存储的转换将被忽略。

transform.merge(变换)

({
   transforms = table  --
})

该函数需要 table 的转换 将它们合并为一个转换。 一旦应用于输入,此转换将产生 table 的输出, 包含转换后的输入。

例如,以下代码:

> f = transform.merge{
        [1] = function(x) return 2*x end,
        [2] = function(x) return x + 10 end,
        foo = function(x) return x / 2 end,
        [4] = function(x) return x - x end
   }

生成一个函数,该函数将一组转换应用于同一输入:

> f(3)
   {
     1 : 6
     2 : 13
     foo : 1.5
     4 : 0
   }

transform.tablenew()

该函数根据一个函数创建一个新的函数表 现有的函数表。

transform.tableapply(变换)

({
   transform = function  --
})

此函数对输入表应用转换。 它返回与输入大小相同的输出表。

例如,以下代码:

> f = transform.tableapply(function(x) return 2*x end)

生成一个将任何输入乘以 2 的函数:

> f({[1] = 1, [2] = 2, foo = 3, [4] = 4})
   {
     1 : 2
     2 : 4
     foo : 6
     4 : 8
   }

transform.tablemergekeys()

该函数按键合并表。更准确地说,输入必须是 tabletable,此函数将反转该表以便 make the keys from the nested table accessible first.

例如,如果输入是:

> x = { sample1 = {input = 1, target = "a"} , sample2 = {input = 2, target = "b", flag = "hard"}

然后应用这个函数将产生:

> transform.tablemergekeys(x)
{
   input :
         {
           sample1 : 1
           sample2 : 2
         }
   target :
          {
            sample1 : "a"
            sample2 : "b"
          }
   flag :
        {
           sample2: "hard"
        }
}

transform.makebatch([合并])

({
  [merge = function]  --
})

很多tnt.Dataset都用这个函数来格式化 样本采用 tnt.Engine 使用的格式。

该函数首先将合并密钥到 产生一个输出表。然后,将该表转换为张量: 使用用户提供的 merge 转换或 只需直接将表连接成张量即可。

该函数使用 组成 变换来应用 连续的转变。

transform.randperm(尺寸)

({
   size = number  --
})

此函数创建一个向量,其中包含从 1 到 size 的索引排列。 该向量是 LongTensor 并且 size 必须是数字。

创建向量后,该函数可用于调用其中的特定索引。

例如:

> p = transform.randperm(3)

创建一个包含索引排列的函数 p

> p(1)
2
> p(2)
1
> p(3)
3

transform.normalize([阈值])

({
  [threshold = number]  --  [default=0]
})

此函数对数据 i.e 进行标准化,删除其平均值和 将其除以标准差。

输入必须是 Tensor

创建后,可以给出threshold(必须是数字)。然后, 数据将除以标准差,前提是 偏差大于threshold。这很方便,如果 偏差很小,除以它可能会导致不稳定。

tnt.ListDataset(自身,列表,加载[,路径])

({
   self = tnt.ListDataset  --
   list = tds.Hash         --
   load = function         --
  [path = string]          --
})

考虑 list(可以是 tds.Hashtabletorch.LongTensor) 数据集的第 i 个样本将由 load(list[i]) 返回,其中 load() 是 用户提供的闭包。

如果提供了 path,则列表被假定为字符串列表,并且将 当输入到 load() 时,每个元素 list[i] 都会以 path/ 为前缀。

目的:许多中低规模的数据集可以看作文件列表 (例如表示输入样本)。对于此文件列表,目标 通常可以通过简单的方式推断出来。

tnt.ListDataset(自身, 文件名, 负载[, 最大负载][, 路径])

({
   self     = tnt.ListDataset  --
   filename = string           --
   load     = function         --
  [maxload  = number]          --
  [path     = string]          --
})

filename 指定的文件被解释为字符串列表(一个 每行字符串)。数据集的第 i 个样本将通过以下方式返回 load(line[i]),其中load()是用户提供的闭包 line[i]filename 的 i 系列。

如果提供了 path,则列表被假定为字符串列表,并且将 当输入到 load() 时,每个元素 list[i] 都会以 path/ 为前缀。

tnt.TableDataset(自身,数据)

{
   self = tnt.TableDataset  --
   data = table             --
}

tnt.TableDataset 接口现有数据 到火炬网。如果您想在小型数据集上使用 torchnet,它会很有用。

数据必须包含在 tds.Hash 中。

tnt.TableDataset 对数据进行浅表复制。

构建 tnt.TableDataset 时加载数据:

> a = tnt.TableDataset{data = {1,2,3}}
> print(a:size())
3

tnt.TableDataset 假设表具有从 1 开始的连续键。

tnt.IndexedDataset(self, 字段[, 路径][, maxload][, mmap][, mmapidx][, 独立])

{
   self       = tnt.IndexedDataset  --
   fields     = table               --
  [path       = string]             --
  [maxload    = number]             --
  [mmap       = boolean]            --  [default=false]
  [mmapidx    = boolean]            --  [default=false]
  [standalone = boolean]            --  [default=false]
}

tnt.IndexedDataset() 是一个基于(可能是多个)构建的数据结构 包含一堆相同类型的张量的数据档案。

参见 tnt.IndexedDatasetWriter 和 tnt.IndexedDatasetReader 查看如何创建和 读取单个档案。

目的:大型数据集(包含大量文件)通常不太好 由文件系统处理(尤其是通过网络)。 tnt.IndexedDataset 提供了一种方便有效的方法将它们捆绑到一个单一的 归档文件,与索引文件关联。

如果提供了 path,则 fields 必须是 Lua 数组(键为 数字),其中值是表示文件名前缀的字符串 (索引,存档)对。 换句话说,path/field.{idx,bin} 必须存在。的 该数据集返回的第 i 个样本将是一个包含每个字段的表 作为键,以及在索引 i 处相应档案中找到的张量。

如果未提供 path,则 fields 必须是 Lua 哈希。每个键 代表样本字段,对应的值必须是表格 包含键 idx (对于索引文件名路径)和 bin (对于 存档文件名路径)。

如果提供(且为正),maxload 将数据集大小限制为 指定尺寸。

档案和/或索引也可以使用 mmap 进行内存映射和 mmapidx 标志。

如果 standalone 为 true,则构造函数期望只有一个字段 提供。数据集返回的第 i 个样本将是在以下位置找到的项目 索引 i 处的档案。这对于 table 档案特别有用。

tnt.IndexedDatasetWriter(自身,索引文件名,数据文件名,类型)

({
   self          = tnt.IndexedDatasetWriter  --
   indexfilename = string                    --
   datafilename  = string                    --
   type          = string                    --
})

创建(存档,索引)文件对。存档将包含相同指定 type 的张量。

type 必须是在 {bytecharshortintlongfloatdouble 中选择的字符串或 table}。

indexfilename 是要创建的索引文件的完整路径。 datafilename 是要创建的数据归档文件的完整路径。

使用 add() 将张量添加到存档中。

请注意,您必须调用 close() 以确保所有 数据写入磁盘并创建索引文件。

table 类型比较特殊:数据将存储到 CharTensor 中, 从 Lua 表对象序列化。 IndexedDatasetReader 然后将 在读取时将 CharTensor 反序列化到表中。这允许存储 异构数据轻松导入IndexedDataset。

tnt.IndexedDatasetWriter(自身,索引文件名,数据文件名)

({
   self          = tnt.IndexedDatasetWriter  --
   indexfilename = string                    --
   datafilename  = string                    --
})

打开现有的(存档、索引)文件对以进行追加。张量类型是从提供的推断出来的 索引文件。

indexfilename 是要打开的索引文件的完整路径。 datafilename 是要打开的数据归档文件的完整路径。

tnt.IndexedDatasetWriter.add(自身,张量)

({
   self   = tnt.IndexedDatasetWriter  --
   tensor = torch.*Tensor             --
})

将张量添加到存档中并记录其索引位置。张量类型必须相同 比创建 tnt.IndexedDatasetWriter 时指定的值要高。

tnt.IndexedDatasetWriter.add(自身,文件名)

({
   self     = tnt.IndexedDatasetWriter  --
   filename = string                    --
})

给出一个 filename 的便捷方法将打开相应的 文件以 binary 模式,并读取其中的所有数据,就好像它是该类型一样 在 tnt.IndexedDatasetWriter 构造中指定。 对应的一个 然后将张量添加到 archive/index 对中。

tnt.IndexedDatasetWriter.add(自身,表)

(
   self  = tnt.IndexedDatasetWriter  --
   table = table                     --
)

便捷方法仅适用于 table 类型 IndexedDataset。 该表将被序列化为 CharTensor。

tnt.IndexedDatasetWriter.add(自己)

({
   self = tnt.IndexedDatasetWriter  --
})

完成索引,并关闭 archive/index 文件名对。这个方法 必须调用以确保索引已写入并且所有归档数据均已写入 刷新到磁盘上。

tnt.IndexedDatasetReader(self, 索引文件名, 数据文件名[, mmap][, mmapidx])

({
   self          = tnt.IndexedDatasetReader  --
   indexfilename = string                    --
   datafilename  = string                    --
  [mmap          = boolean]                  --  [default=false]
  [mmapidx       = boolean]                  --  [default=false]
})

读取之前创建的 archive/index 对 tnt.IndexedDatasetWriter.

indexfilename 是索引文件的完整路径。 datafilename 是存档文件的完整路径。

可以通过以下方式为存档和索引指定内存映射 可选的 mmapmmapidx 标志。

tnt.IndexedDatasetReader.尺寸(自)

返回存档中存在的张量数量。

tnt.IndexedDatasetReader.get(自身,索引)

返回存档中指定 index 处的张量。

tnt.TransformDataset(自身,数据集,变换[,键])

({
   self      = tnt.TransformDataset  --
   dataset   = tnt.Dataset           --
   transform = function              --
  [key       = string]               --
})

给定一个闭包 transform() 和一个 datasettnt.TransformDataset 当查询样本时以即时的方式应用闭包 tnt.Dataset:get().

如果提供了 key,则闭包将应用于指定的示例字段 由 key(仅)。闭包必须返回新的相应字段值。

如果未提供密钥,则封闭将应用于整个样本。的 闭包必须返回新的样本表。

新数据集的大小等于底层 dataset 的大小。

目的:在进行预处理操作时,方便 能够执行即时转换 数据集。

tnt.TransformDataset(自身、数据集、转换)

({
   self       = tnt.TransformDataset  --
   dataset    = tnt.Dataset           --
   transforms = table                 --
})

给定一组闭包和 datasettnt.TransformDataset 适用 当查询样本时,这些闭包会以即时的方式进行 tnt.Dataset:get().

闭包在 Lua 表 transforms 中提供,其中 (key,value) 对代表一个(示例字段名称,要应用的相应闭包 到字段名称)。

每个闭包必须返回相应字段的新值。

tnt.BatchDataset(自身、数据集、batchsize[、perm][、合并][、策略][、过滤器])

({
   self      = tnt.BatchDataset  --
   dataset   = tnt.Dataset       --
   batchsize = number            --
  [perm      = function]         --  [has default value]
  [merge     = function]         --
  [policy    = string]           --  [default=include-last]
  [filter    = function]         --  [has default value]
})

给定 datasettnt.BatchDataset 将此数据集中的样本合并到 形成一个新样本,可以将其解释为一个批次(大小 batchsize).

merge 函数控制批处理的执行方式。这是一个闭包 将包含所有出现次数的 Lua 数组作为输入(对于给定批次) 样本字段的值,并返回这些字段的聚合版本 发生。默认情况下,出现的次数应该是张量,并且 它们沿着第一维度聚集。

更正式地说,如果基础数据集的第 i 个样本写为:

{input=<input_i>, target=<target_i>}

假设样本中只有两个字段 inputtarget,则 merge() 将传递以下形式的表:

{<input_i_1>, <input_i_2>, ... <input_i_n>}

{<target_i_1>, <target_i_2>, ... <target_i_n>}

n 是批量大小。

在执行批处理时打乱示例通常很重要 操作。 perm(idx, size) 是一个返回混洗索引的闭包 基础数据集中位置 idx 处的样本。为了方便起见, 底层数据集的 size 也传递给闭包。由 默认情况下,闭包是身份。

基础数据集大小可能或可能不总是被整除 batchsize。 可选的 policy 字符串指定如何处理角点 案例:

  • include-last 确保底层数据集的所有样本都能被看到,批次的大小等于或小于 batchsize
  • 如果基础数据集的大小无法正确整除,skip-last 将跳过基础数据集的最后一个示例。批次的大小始终等于 batchsize
  • 如果基础数据集的大小不能被 batchsize 整除,divisible-only 将引发错误。

目的:批次的概念取决于问题。在 torchnet 中,它已启动 供用户将样品解释为批次或非批次。当一个人想要 将现有数据集中的样本组装成一批,然后 tnt.BatchDataset 适合这项工作。有时却更多 方便从头开始编写数据集,提供“批量”样本。

tnt.CoroutineBatchDataset(自身、数据集、batchsize[、perm][、合并][、策略][、过滤器])

({
   self      = tnt.CoroutineBatchDataset  --
   dataset   = tnt.Dataset                --
   batchsize = number                     --
  [perm      = function]                  --  [has default value]
  [merge     = function]                  --
  [policy    = string]                    --  [default=include-last]
  [filter    = function]                  --  [has default value]
})

给定 datasettnt.CoroutineBatchDataset 合并来自该数据集的样本 形成一个新样本,该样本可以解释为一个批次(大小为 batchsize)。

它的行为与 tnt.BatchDataset 相同并且具有相同的参数(请参阅 文档中提供了更多详细信息),但有一个重要区别: 它允许底层数据集推迟返回单个样本 一次通过调用 coroutine.yield() (来自底层数据集)。

当需要使用低效或缓慢的数据集时,这非常有用 致电 dataset:get() 后立即提供所需样品。的 底层 dataset:get() 中的一般代码模式为:

FooDataset.get = function(self, idx)
   prepare(idx)  -- stores sample in self.__data[idx]
   coroutine.yield()
   return self.__data[idx]
end

这里,函数 prepare(idx) 可以实现,例如,缓冲 在实际获取索引之前。

tnt.ConcatDataset(自身,数据集)

{
   self     = tnt.ConcatDataset  --
   datasets = table              --
}

给定一个 Lua 数组 (datasets) tnt.Dataset,连接 将它们合并到一个数据集中。 新数据集的大小是以下数据集的总和 基础数据集大小。

目的:可能有助于组装不同的现有数据集 大型数据集,因为串联操作是在 即时方式。

tnt.ResampleDataset(自身,数据集[,采样器] [,大小])

给定 dataset,创建一个新的数据集,该数据集将从该数据集(重新)采样 使用提供的 sampler(dataset, idx) 闭包的底层数据集。

如果提供了 size,则新创建的数据集将具有 指定 size,这可能与底层数据集不同 尺寸。

如果未提供 size,则新数据集将具有相同的大小 比底层的。

默认情况下,sampler(dataset, idx) 是身份,简单来说就是 returning idxdataset对应于构建时提供的底层数据集,并且 idx 可以取 1 到 size 之间的值。它必须返回范围内的索引 对于底层数据集来说是可接受的。

目的:打乱数据、重新加权样本、获取样本的子集 数据。请注意,一个重要的子类是(tnt.ShuffleDataset), 为方便起见而提供。

tnt.ShuffleDataset(自身,数据集[,大小][,替换])

({
   self        = tnt.ShuffleDataset  --
   dataset     = tnt.Dataset         --
  [size        = number]             --
  [replacement = boolean]            --  [default=false]
})

tnt.ShuffleDataset 是以下子类 tnt.ResampleDataset 是为了方便起见。

它从给定的 dataset 中均匀采样,有或没有 replacement。可以通过调用重新绘制所选分区 重新采样()。

如果replacementtrue,那么指定的size可能大于 底层 dataset

如果未提供 size,则新数据集大小将等于 底层 dataset 大小。

目的:最简单的打乱数据集的方法!

tnt.ShuffleDataset.重采样(自身)

tnt.ShuffleDataset 相关的排列是固定的,这样两个 对相同索引的调用将从底层返回相同的样本 数据集。

调用 resample() 随机抽取一个新的排列。

tnt.SplitDataset(自身,数据集,分区[,初始分区])

({
   self             = tnt.SplitDataset  --
   dataset          = tnt.Dataset       --
   partitions       = table             --
  [initialpartition = string]           --
})

根据指定的partitions,对给定的dataset进行分区。 使用 方法 select() 选择当前分区 在使用中。

Lua 哈希表 partitions 的形式为 (key, value),其中 key 是 用户选择的字符串命名分区,值是代表的数字 重量(0 到 1 之间的数字)或大小(样本数量) 对应的分区。

分区是线性实现的(无混洗)。参见 tnt.ShuffleDataset 如果你想打乱数据集 分区之前。

可选变量initialpartition指定加载的分区 最初。

目的:在机器学习中用于执行验证程序。

tnt.SplitDataset.select(自身,分区)

({
   self      = tnt.SplitDataset  --
   partition = string            --
})

将当前使用的分区切换到partition指定的分区, 它必须是与以下位置提供的名称之一相对应的字符串 建设。

当前数据集大小以及返回的样本都会相应变化 通过 get() 方法。

数据集迭代器

使用 for 循环可以轻松迭代数据集。然而,有时 人们想要以一种即时的方式或线程样本获取的方式过滤掉样本。

迭代器适用于这种特殊情况。一般来说,不要写 用于处理自定义情况的迭代器,并改为编写 tnt.Dataset

迭代器实现两种方法:

  • run() 返回一个可在 for 循环中使用的 Lua 迭代器。
  • exec(funcname, ...) 在底层数据集上执行给定的 funcname。

典型用法是通过 for 循环实现的:

for sample in iterator:run() do
  <do something with sample>
end

迭代器实现 __call 事件,因此也可以使用 () 运算符:

for sample in iterator() do
  <do something with sample>
end

tnt.DatasetIterator(自身, 数据集[, 烫发][, 过滤器][, 变换])

({
   self      = tnt.DatasetIterator  --
   dataset   = tnt.Dataset          --
  [perm      = function]            --  [has default value]
  [filter    = function]            --  [has default value]
  [transform = function]            --  [has default value]
})

默认数据集迭代器。

perm(idx) 是用于打乱示例的排列。如果洗牌 需要的话,可以使用这个闭包,或者(更好)使用 基础数据集上的 tnt.ShuffleDataset。

filter(sample) 是一个闭包,如果给定样本则返回 true 应考虑或 false 如果没有。

transform(sample)是一个闭包,可以执行在线转换 样品。它返回给定 sample 的修改版本。它是 默认身份。使用起来往往更有趣 tnt.TransformDataset 用于此目的。

tnt.DatasetIterator.exec(tnt.DatasetIterator,名称,...)

在底层数据集上执行给定的方法 name,并将其传递给 后续参数,并返回 name 方法返回的内容。

tnt.ParallelDatasetIterator(self[, init], 闭包, nthread[, perm][, 过滤器][, 变换][, 有序])

({
   self      = tnt.ParallelDatasetIterator  --
  [init      = function]                    --  [has default value]
   closure   = function                     --
   nthread   = number                       --
  [perm      = function]                    --  [has default value]
  [filter    = function]                    --  [has default value]
  [transform = function]                    --  [has default value]
  [ordered   = boolean]                     --  [default=false]
})

允许在线程中迭代数据集 方式。 tnt.ParallelDatasetIterator:run()保证所有样品 会被看到,但不保证顺序,除非 ordered 设置为 true。

此类的目的是实现零预处理成本。 当从以下位置即时读取数据集时 磁盘(未将它们完全加载到内存中),或执行复杂的 预处理这可能很有趣。

用于并行化的线程数由 nthread 指定。

init(threadid)(其中 threadid=1..nthread)是一个闭包,可以 如果需要的话,根据需要初始化指定的线程。它什么也没做 默认情况下。

closure(threadid) 将在每个线程上调用,并且必须返回 tnt.Dataset 实例。

perm(idx) 是用于打乱示例的排列。如果洗牌是 需要的话,可以使用这个闭包,或者(更好)使用 基础数据集上的 tnt.ShuffleDataset (由 closure() 返回)。

filter(sample) 是一个闭包,如果给定样本则返回 true 应考虑或 false 如果没有。请注意,过滤器被称为_after_ 以线程方式获取数据。

transform(sample) 是将给定样本映射到新值的函数。 这种转换发生在过滤之前。

ordered 设置为 true 时,迭代器返回的样本的顺序 是有保证的。此选项对于可重复的实验特别有用。 默认情况下 ordered 为 false,这意味着顺序不被保证 run()(尽管实践中顺序通常相似)。

此数据集引发的一个常见错误是 closure() 不是 可序列化。确保 closure() 的所有 上值 均为 可序列化。建议不惜一切代价避免 升值, 并确保您需要(反)序列化所需的所有适当的火炬包 init() 函数中的 closure()

有关更多信息,请查看 线程包, tnt.ParallelDatasetIterator 所依赖的。

tnt.ParallelDatasetIterator.execSingle(tnt.DatasetIterator,名称,...)

在第一个对应的数据集上执行给定的方法 name 可用线程,向其传递后续参数,并返回 name 方法返回。

例如:

  local iterator = tnt.ParallelDatasetIterator{...}
  print(iterator:execSingle("size"))

将打印第一个可用线程中加载的数据集的大小。

tnt.ParallelDatasetIterator.exec(tnt.DatasetIterator,名称,...)

在每个线程中的底层数据集上执行给定的方法 name, 将后续参数传递给每个人,并返回一个表 name 方法为每个线程返回的内容。

例如:

  local iterator = tnt.ParallelDatasetIterator{...}
  for _, v in pairs(iterator:exec("size")) do
      print(v)
  end

将打印每个线程中加载的数据集的大小。

tnt.Engine

在尝试不同的模型和数据集时,底层训练 程序通常是相同的。 Engine 模块提供样板逻辑 模型训练和测试所必需的。这可能包括进行 模型(nn.Module)、tnt.DatasetIterators、 nn.Criterions 和 tnt.Meters。

tnt.Engine() 的实例 engine 实现了两个主要方法:

  • engine:train(),用于训练数据模型 (i.e. sample data, forward prop, backward prop).
  • engine:test(),用于评估数据模型 (optionally with respect to a nn.Criterion).

Engine可以实现任何常见的底层训练和测试 涉及模型和数据的过程。它还可以设计为允许用户 某些事件后的控制,例如前向传播、标准评估或 通过使用协程来结束纪元(参见 tnt.SGDEngine)。

tnt.SGDEngine

SGDEngine模块实现随机梯度下降训练 train中的过程,包括数据采样、前向传播、后向传播和 参数更新。它还作为协程运行,允许用户控制 (i.e。在“开始”等事件中增加某种 tnt.Meter), “开始纪元”、“向前”、“向前标准”、“向后”等。 可用的钩子如下:

hooks = {
   ['onStart']             = function() end, -- Right before training
   ['onStartEpoch']        = function() end, -- Before new epoch
   ['onSample']            = function() end, -- After getting a sample
   ['onForward']           = function() end, -- After model:forward
   ['onForwardCriterion']  = function() end, -- After criterion:forward
   ['onBackwardCriterion'] = function() end, -- After criterion:backward
   ['onBackward']          = function() end, -- After model:backward
   ['onUpdate']            = function() end, -- After UpdateParameters
   ['onEndEpoch']          = function() end, -- Right before completing epoch
   ['onEnd']               = function() end, -- After training
}

要为给定的钩子指定新的闭包,我们可以使用以下命令访问它 engine.hooks.<onEvent>。例如,我们可以在每次之前重置 Meter 纪元:

local engine = tnt.SGDEngine()
local meter  = tnt.AverageValueMeter()
engine.hooks.onStartEpoch = function(state)
   meter:reset()
end

因此,train 需要一个网络(nn.Module),这是一个表达 损失函数(nn.Criterion)、数据集迭代器(tnt.DatasetIterator)和 学习率,至少。 test 功能可进行简单评估 数据集上的模型。

维护 state 用于外部访问模块的输出和参数 以及采样数据。 state表的内容如下,其中 传递的值来自 engine:train() 的参数:

state = {
   ['network']     = network,
   ['criterion']   = criterion,
   ['iterator']    = iterator,
   ['lr']          = lr,
   ['lrcriterion'] = lrcriterion,
   ['maxepoch']    = maxepoch,
   ['sample']      = {},
   ['epoch']       = 0, -- epoch done so far
   ['t']           = 0, -- samples seen so far
   ['training']    = true
}

tnt.OptimEngine

OptimEngine 模块封装了来自 https://github.com/torch/optim. 训练开始时,引擎会调用 getParameters 在提供的网络上。

train 方法除了需要以下参数之外,还需要以下参数 SGDEngine.train参数:

  • optimMethod优化函数(e.g optim.sgd
  • config 包含优化器配置参数的表

示例:

  local engine = tnt.OptimEngine()
  engine:train{
     network = model,
     criterion = criterion,
     iterator = iterator,
     optimMethod = optim.sgd,
     config = {
        learningRate = 0.1,
        momentum = 0.9,
     },
  }

tnt.Meter

训练模型时,您通常希望测量模型的表现如何 表演。具体来说,您可能想要测量平均处理时间 每批数据所需的分类器 a 的分类错误或 AUC 验证集,或检索模型的 precision@k。

仪表提供了一种标准化的方法来测量一系列不同的措施, 这使得测量模型的各种属性变得容易。

几乎所有仪表(tnt.TimeMeter 除外)都实现三种方法:

  • add() 向仪表添加观察。
  • value() 返回仪表的值,考虑所有观测值。
  • reset() 删除所有先前添加的观测值,重置仪表。

add() 方法的确切输入参数因仪表而异。 大多数仪表将方法定义为 add(output, target),其中 output 是 模型产生的输出,target 是数据的真实标签。

value() 方法对于大多数仪表来说是无参数的,但对于以下测量: 有一个参数(例如 precision@k 中的 k 参数),它们可能需要一个 输入参数。

仪表的典型用法示例如下:

local meter = tnt.<Measure>Meter()  -- initialize meter
for state, event in tnt.<Optimization>Engine:train{
   network   = network,
   criterion = criterion,
   iterator  = iterator,
} do
  if state == 'start-epoch' then
     meter:reset()  -- reset meter
  elseif state == 'forward-criterion' then
     meter:add(state.network.output, sample.target)  -- add value to meter
  elseif state == 'end-epoch' then
     print('value of meter:' .. meter:value())  -- get value of meter
  end
end

tnt.APMeter(自)

({
   self = tnt.APMeter  --
})

tnt.APMeter 测量每个类别的平均精度。

tnt.APMeter 设计用于在 NxK 张量 outputtarget,以及可选的 Nx1 张量权重,其中 (1) output 包含 N 示例和 K 类的模型输出分数应该更高 当模型更加确信该示例应该被正面标记时, 当模型认为该示例应该被负面标记时更小 (例如,sigmoid 函数的输出); (2) target 包含 仅值 0(对于负例)和 1(对于正例);和(3) weight ( > 0) 表示每个样本的重量。

tnt.APMeter无参数需要设置。

tnt.AverageValueMeter(自)

({
   self = tnt.AverageValueMeter  --
})

tnt.AverageValueMeter 测量并返回平均值和 任何 added 的数字集合的标准差。它是 例如,可用于测量一组示例的平均损失。

add() 函数期望输入 Lua 数字 value,即值 需要将其添加到要平均的值列表中。它还作为输入 一个可选参数 n,为平均值中的 value 分配一个权重,在 为了便于计算加权平均值(默认 = 1)。

tnt.AverageValueMeter 初始化时没有需要设置的参数。

tnt.AUCMeter(自)

({
   self = tnt.AUCMeter  --
})

tnt.AUCMeter 测量接收器工作特性下的面积 (ROC) 二元分类问题的曲线。曲线下面积 (AUC) 可以解释为给定随机选择的正数的概率 例子和随机选择的负例,正例是 分类模型赋予比负例更高的分数。

tnt.AUCMeter 设计用于在一维张量 output 上运行 和 target,其中 (1) output 包含应该 当模型更加确信该示例应该是积极的时,该值会更高 标记,并且当模型认为该示例应该为负时较小 带标签(例如 sigmoid 函数的输出); (2) target 仅包含值 0(对于负示例)和 1(对于正示例)。

tnt.AUCMeter无参数需要设置。

tnt.ConfusionMeter(自身, k[, 归一化])

{
   self       = tnt.ConfusionMeter  --
   k          = number              --
  [normalized = boolean]            --  [default=false]
}

tnt.ConfusionMeter 为多类构造混淆矩阵 分类问题。它不支持多标签、多类问题: 对于此类问题,请使用tnt.MultiLabelConfusionMeter

初始化时,参数k表示数量 必须指定所考虑的分类问题中的类别。 此外,可选参数 normalized(默认 = false)可以是 指定确定混淆矩阵是否标准化 (即,它包含百分比)或不包含(即,它包含计数)。

add(output, target) 方法将 NxK 张量 output 作为输入, 包含从模型中获得的 N 个示例和 K 个类别的输出分数, 以及提供目标的相应 N 张量或 NxK-tensor target 对于N个例子。当 target 是 N 张量时,假设目标是 1 到 K 之间的整数值。当目标是 NxK-tensor 时,目标是 假定作为 one-hot 向量提供(即,仅包含 0 和要编码的目标值位置处的单个 1)。

value() 方法没有参数,并以 a 形式返回混淆矩阵 KxK 张量。在混淆矩阵中,行对应于地面实况目标, 列对应于预测目标。

tnt.mAPMeter(自)

({
   self = tnt.mAPMeter  --
})

tnt.mAPMeter 测量所有类别的平均精度。

tnt.mAPMeter 设计用于在 NxK 张量 outputtarget,以及可选的 Nx1 张量权重,其中 (1) output 包含 N 示例和 K 类的模型输出分数应该更高 当模型更加确信该示例应该被正面标记时, 当模型认为该示例应该被负面标记时更小 (例如,sigmoid 函数的输出); (2) target 包含 仅值 0(对于负例)和 1(对于正例);和(3) weight (> 0) 表示每个样本的重量。

tnt.mAPMeter无参数需要设置。

tnt.MovingAverageValueMeter(自身,窗口大小)

({
   self       = tnt.MovingAverageValueMeter  --
   windowsize = number                       --
})

tnt.MovingAverageValueMeter 测量并返回平均值 以及任何 added 数字集合的标准差 在最近的移动平均线窗口内。它很有用,例如, 衡量一组样本的平均损失 最近的窗口。

add() 函数期望输入 Lua 数字 value,即值 需要将其添加到要平均的值列表中。

tnt.MovingAverageValueMeter 需要将移动窗口大小设置为 初始化时间。

tnt.MultiLabelConfusionMeter(自身, k[, 归一化])

{
   self       = tnt.MultiLabelConfusionMeter  --
   k          = number                        --
  [normalized = boolean]                      --  [default=true]
}

tnt.MultiLabelConfusionMeter 构造了一个混淆矩阵 标签,多类分类问题。在构建混乱的过程中 矩阵,假设正预测的数量等于 真实情况中的积极标签。正确的预测(即,标签 也在地面实况集中的预测集)被添加到 混淆矩阵的对角线。不正确的预测(即,标签中 不在真实数据集中的预测集)均等地划分为 地面实况集中所有非预测标签。

初始化时,参数k表示数量 必须指定所考虑的分类问题中的类别。 此外,可选参数 normalized(默认 = false)可以是 指定确定混淆矩阵是否标准化 (即,它包含百分比)或不包含(即,它包含计数)。

add(output, target) 方法将 NxK 张量 output 作为输入, 包含从模型中获得的 N 个示例和 K 个类别的输出分数, 以及相应的 NxK-tensor target 为 N 提供目标 使用 one-hot 向量(即仅包含零和一个的向量)的示例 待编码目标值位置处的单个)。

value() 方法没有参数,并以 a 形式返回混淆矩阵 KxK 张量。在混淆矩阵中,行对应于地面实况目标, 列对应于预测目标。

tnt.ClassErrorMeter(self[, topk][, 准确度])

{
   self     = tnt.ClassErrorMeter  --
  [topk     = table]               --  [has default value]
  [accuracy = boolean]             --  [default=false]
}

tnt.ClassErrorMeter 测量分类误差(以 % 为单位) 分类模型(零一损失)。该仪表还可以测量误差 预测前 k 个评分标签中的正确标签(例如,在 Imagenet 竞赛,通常测量分类@5 错误)。

初始化时,它需要可选参数:(1)一个表 topk 包含分类@k 错误应达到的值 措施(默认 = {1}); (2) 布尔值 accuracy 使仪表 输出精度而不是误差(精度 = 1 - 误差)。

add(output, target) 方法将 NxK-tensor output 作为输入, 包含 N 个示例和 K 个类别中每个示例的输出分数, 和一个 N 张量 target,其中包含与每个 N 个示例(目标是 1 到 K 之间的整数)。如果只有一个例子 added、output 也可以是 K 张量并以 1 张量为目标。

请注意,topk(如果指定)不得包含大于 K 的值。

value() 返回一个表,其中所有值的分类@k 错误 在初始化时在 topk 中指定的 k 处。或者, value(k) 以数字形式返回分类@k 错误;仅 k 的值 topk 的元素是允许的。如果 accuracy 设置为 true 初始化时,value() 方法返回精度而不是错误。

tnt.TimeMeter(自身[,单位])

({
   self = tnt.TimeMeter  --
  [unit = boolean]       --  [default=false]
})

tnt.TimeMeter 旨在测量事件之间的时间,并且可以 例如,用于测量每批数据的平均处理时间。 它与大多数其他仪表的不同之处在于它提供的方法:

在初始化时,可以提供可选的布尔参数 unit (默认 = false)。当设置为true时,仪表返回的值 将除以 incUnit() 方法被调用的次数。 例如,这允许用户计算每次的平均处理时间 处理批处理后,只需调用 incUnit() 方法即可进行批处理。

tnt.TimeMeter提供以下方法:

  • reset() 重置定时器,将定时器和单位计数器设置为零。
  • stop() 停止定时器。
  • resume() 恢复定时器。
  • incUnit() 将单位计数器加一。
  • value() 返回自上次 reset() 以来经过的时间;除以 unit=true 时的计数器值。

tnt.PrecisionAtKMeter(self[, topk][, 暗淡][, 在线])

{
   self   = tnt.PrecisionAtKMeter  --
  [topk   = table]                 --  [has default value]
  [dim    = number]                --  [default=2]
  [online = boolean]               --  [default=false]
}

tnt.PrecisionAtKMeter 测量预先指定的排序方法的精度@k 水平 k. precision@k 是排名前 k 的百分比 根据正确(正)目标列表中的模型的项目。

在初始化时,可以给出一个表 topk 作为指定的输入 将测量 precision@k 的级别 k(默认 = {10})。在 另外,可以提供数字dim来指定在哪个维度上 应该计算 precision@k (默认 = 2),并且布尔值 online 可能是 指定指示我们是否一次看到沿维度 dim 的所有输入 (默认 = false)。

add(output, target) 方法采用两个输入。默认模式下(dim=2online=false),输入意味着:

  • 一个 NxC 张量,对于 N 个示例(查询)中的每个示例都包含一个分数 indicating to what extent each of the C classes (documents) is relevant to the query, according to the model.
  • 二进制 NxC target 张量,编码 C 类中的哪一个 (documents) are actually relevant to the the N-th input (query). For instance, a row of {0, 1, 0, 1} indicates that the example is associated with classes 2 and 4.

dim 设置为 1 的结果与转置张量相同 上面的outputtarget。设置online=true的结果是 该函数假设不是查询数量 N 正在增长 多次调用add(),但候选文件数量为C。 (使用 当C较大而N较小的场景时使用此模式。)

value() 方法返回一个表,其中包含 precision@k(即 正确预测的目标的百分比)在 topk 的截止水平 在初始化时指定。或者,精度@k位于 可以通过调用value(k)来获得特定的级别k。注意级别 k 应该是初始化时指定的表 topk 的元素。

请注意,topk 中的最大值不能高于总和 类(文档)的数量。

tnt.RecallMeter(自身[,阈值][,每类])

{
   self      = tnt.RecallMeter  --
  [threshold = table]           --  [has default value]
  [perclass  = boolean]         --  [default=false]
}

tnt.RecallMeter 测量预排序方法的召回率 指定的阈值。召回率是正确(正)的百分比 根据模型位于正标记项目列表中的目标。

在初始化时,tnt.RecallMeter 提供两个可选的 参数。第一个参数是一个表threshold,其中包含所有 测量召回率的阈值(默认 = {0.5})。阈值 应该是 0 到 1 之间的数字。第二个参数是布尔值 perclass 当设置为 true 时,仪表会测量每类的召回率 (默认 = false)。当perclass设置为false时,召回很简单 对所有示例进行平均。

add(output, target) 方法采用两个输入:

  • 一个 NxK 张量,对于 N 个示例中的每个示例指示概率 of the example belonging to each of the K classes, according to the model. The probabilities should sum to one over all classes; that is, the row sums of output should all be one.
  • 二进制 NxK target 张量,编码 K 类中的哪一个 are associated with the N-th input. For instance, a row of {0, 1, 0, 1} indicates that the example is associated with classes 2 and 4.

value() 方法返回一个包含模型召回率的表 在初始化时指定的 thresholds 处测量的预测。的 value(t) 方法返回特定阈值 t 的召回率。请注意 该阈值 t 应该是在以下位置指定的 threshold 表的元素 仪表的初始化时间。

tnt.PrecisionMeter(自身[,阈值][,每类])

{
   self      = tnt.PrecisionMeter  --
  [threshold = table]              --  [has default value]
  [perclass  = boolean]            --  [default=false]
}

tnt.PrecisionMeter 测量预排序方法的精度 指定的阈值。精度是阳性标记的百分比 根据正确(正)目标列表中的模型的项目。

在初始化时,tnt.PrecisionMeter 提供两个可选的 参数。第一个参数是一个表threshold,其中包含所有 测量精度的阈值(默认 = {0.5})。阈值 应该是 0 到 1 之间的数字。第二个参数是布尔值 perclass 当设置为 true 时,使仪表测量每级的精度 (默认 = false)。当 perclass 设置为 false 时,精度为 对所有示例进行平均。

add(output, target) 方法采用两个输入:

  • 一个 NxK 张量,对于 N 个示例中的每个示例指示概率 of the example belonging to each of the K classes, according to the model. The probabilities should sum to one over all classes; that is, the row sums of output should all be one.
  • 二进制 NxK target 张量,编码 K 类中的哪一个 are associated with the N-th input. For instance, a row of {0, 1, 0, 1} indicates that the example is associated with classes 2 and 4.

value() 方法返回包含模型精度的表 在初始化时指定的 thresholds 处测量的预测。的 value(t) 方法返回特定阈值 t 处的精度。请注意 该阈值 t 应该是在以下位置指定的 threshold 表的元素 仪表的初始化时间。

tnt.NDCGMeter(自身[, K])

{
   self = tnt.NDCGMeter  --
  [K    = table]         --  [has default value]
}

tnt.NDCGMeter 测量标准化贴现累积增益 (NDCG) 由模型在预先指定的级别 k 生成的排名,并对 NDCG 进行平均 超过所有的例子。

k 级贴现累积增益定义为:

DCG_k = rel_1 + sum{i = 2}^k (rel_i / log_2(i))

这里,rel_i是由外部评估者指定的项目i的相关性。 对于给定示例,将理想的 DCG (IDCG) 定义为最佳可能的 DCG,即 NDCG k 级定义为:

NDCG_k = DCG_k / IDCG_k

在初始化时,仪表将表 K 作为输入,其中包含所有 计算 NDCG 的级别 k。

add(output, relevance) 方法将模型的 NxC 张量作为输入 (1) outputs,对一批 N 个示例的所有 C 个可能输出进行评分; (2) NxC 张量 relevance 包含相应的相关性 这些分数由外部评估者提供。相关性一般是 从人类评估者那里获得。

value() 方法返回一个表,其中包含所有 NDCG 值 初始化时提供的级别 K。或者,NDCG 位于 可以通过调用value(k)来获得特定的级别k。注意级别 k 应该是初始化时指定的表 K 的元素。

请注意,输出的数量和相关性 C 应始终位于 至少与电表计算的最高 NDCG 级别 k 一样高。

tnt.Log

Log 类充当由字符串键索引的表。允许的键必须是 施工时提供。还可以设置特殊密钥 __status__ 方便方法 log:status() 记录基本消息。

查看器闭包可以附加到 Log,并在不同的事件中调用:

  • onSet(log, key, value):当使用log:set{}设置Log的密钥时。
  • onGet(log, key):使用log:get()查询密钥时。
  • onFlush(log):用log:flush()刷新Log的存储数据时。
  • onClose(log):当用 log:close() 关闭 Log 时。

典型的查看器闭包是 textjson,它们允许写入磁盘 或到控制台 Log 存储的密钥子集,在特定的 格式。特殊观察器闭合件 status 是为了在 set() 上调用而设计的 事件,并且只会打印出状态记录。

一个典型的用例如下:

tnt = require 'torchnet'

-- require the viewers we want
logtext = require 'torchnet.log.view.text'
logstatus = require 'torchnet.log.view.status'

log = tnt.Log{
   keys = {"loss", "accuracy"},
   onFlush = {
      -- write out all keys in "log" file
      logtext{filename='log.txt', keys={"loss", "accuracy"}, format={"%10.5f", "%3.2f"}},
      -- write out loss in a standalone file
      logtext{filename='loss.txt', keys={"loss"}},
      -- print on screen too
      logtext{keys={"loss", "accuracy"}},
   },
   onSet = {
      -- add status to log
      logstatus{filename='log.txt'},
      -- print status to screen
      logstatus{},
   }
}

-- set values
log:set{
  loss = 0.1,
  accuracy = 97
}

-- write some info
log:status("hello world")

-- flush out log
log:flush()

tnt.Log(自身,键[,onClose][,onFlush][,onGet][,onSet])

{
   self    = tnt.Log  --
   keys    = table    --
  [onClose = table]   --
  [onFlush = table]   --
  [onGet   = table]   --
  [onSet   = table]   --
}

使用允许的键(字符串)keys 创建新的 Log。 指定事件 带有函数表 onCloseonFlushonGetonSet 的闭包, 当 close()flush()get()set{} 时将被调用 方法将分别被调用。

tnt.Log:状态(自身[,消息][,时间])

({
   self    = tnt.Log   --
  [message = string]   --
  [time    = boolean]  --  [default=true]
})

记录状态消息以及事件的相应(可选)时间。

tnt.Log:设置(自身,密钥)

(
   self = tnt.Log  --
   keys = table    --
)

将多个键(构造时提供的键的子集)设置为 他们对应的值。

将调用附加到 onSet(log, key, value) 事件的闭包。

tnt.Log:获取(自身,密钥)

({
   self = tnt.Log  --
   key  = string   --
})

获取给定键的值。

将调用附加到 onGet(log, key) 事件的闭包。

tnt.Log:齐平(自)

({
   self = tnt.Log  --
})

刷新(清空)日志数据。

将调用附加到 onFlush(log) 事件的闭包。

tnt.Log:关闭(自身)

({
   self = tnt.Log  --
})

关闭日志。

将调用附加到 onClose(log) 事件的闭包。

tnt.Log:附加(自身,事件,闭包)

({
   self     = tnt.Log  --
   event    = string   --
   closures = table    --
})

将一组函数(在表中提供)附加到给定事件。

相关文章

精彩推荐