最新下载
热门教程
- 1
- 2
- 3
- 4
- 5
- 6
- 7
- 8
- 9
- 10
torchnet:实践指南
时间:2026-09-11 17:28:01 编辑:袖梨 来源:一聚教程网
实际评估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
请注意,使用键 foo 和 4 存储的转换将被忽略。
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()
该函数按键合并表。更准确地说,输入必须是
table 或 table,此函数将反转该表以便
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.Hash、table 或 torch.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 必须是在 {byte、char、short、int、long、float、double 中选择的字符串或 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 是存档文件的完整路径。
可以通过以下方式为存档和索引指定内存映射
可选的 mmap 和 mmapidx 标志。
tnt.IndexedDatasetReader.尺寸(自)
返回存档中存在的张量数量。
tnt.IndexedDatasetReader.get(自身,索引)
返回存档中指定 index 处的张量。
tnt.TransformDataset(自身,数据集,变换[,键])
({
self = tnt.TransformDataset --
dataset = tnt.Dataset --
transform = function --
[key = string] --
})
给定一个闭包 transform() 和一个 dataset、tnt.TransformDataset
当查询样本时以即时的方式应用闭包
tnt.Dataset:get().
如果提供了 key,则闭包将应用于指定的示例字段
由 key(仅)。闭包必须返回新的相应字段值。
如果未提供密钥,则封闭将应用于整个样本。的 闭包必须返回新的样本表。
新数据集的大小等于底层 dataset 的大小。
目的:在进行预处理操作时,方便 能够执行即时转换 数据集。
tnt.TransformDataset(自身、数据集、转换)
({
self = tnt.TransformDataset --
dataset = tnt.Dataset --
transforms = table --
})
给定一组闭包和 dataset,tnt.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]
})
给定 dataset,tnt.BatchDataset 将此数据集中的样本合并到
形成一个新样本,可以将其解释为一个批次(大小
batchsize).
merge 函数控制批处理的执行方式。这是一个闭包
将包含所有出现次数的 Lua 数组作为输入(对于给定批次)
样本字段的值,并返回这些字段的聚合版本
发生。默认情况下,出现的次数应该是张量,并且
它们沿着第一维度聚集。
更正式地说,如果基础数据集的第 i 个样本写为:
{input=<input_i>, target=<target_i>}
假设样本中只有两个字段 input 和 target,则 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]
})
给定 dataset,tnt.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 idx。
dataset对应于构建时提供的底层数据集,并且
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。可以通过调用重新绘制所选分区
重新采样()。
如果replacement是true,那么指定的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 ann.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.goptim.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 张量 output 和
target,以及可选的 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 张量 output 和
target,以及可选的 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=2
和 online=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 的结果与转置张量相同
上面的output和target。设置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
outputshould 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
outputshould 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时。
典型的查看器闭包是 text 或 json,它们允许写入磁盘
或到控制台 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。 指定事件
带有函数表 onClose、onFlush、onGet 和 onSet 的闭包,
当 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 --
})
将一组函数(在表中提供)附加到给定事件。
相关文章
- 新版tplink路由器网速慢怎么办(新版tplink路由器网速慢怎么解决) 09-11
- DropDownView:实践指南 09-11
- Onboard-SDK:实践指南 09-11
- itflow:实践指南 09-11
- kazoo:实践指南 09-11
- hexo-theme-shoka:实践指南 09-11