Skip to content

Apache MXNet - KVStore 和可视化

本章探讨 Apache MXNet 的两个关键方面:用于分布式训练和数据同步的 KVStore,以及网络可视化技术。

Key-Value Store (KVStore) 是 MXNet 中的一个基础组件,对于跨多个设备(GPU/CPU)和多台机器训练模型至关重要。它促进了训练过程中模型参数(权重和偏置)的通信和同步。

KVStore 操作的关键概念:

  • KVStore 中共享的每条数据都由一个唯一的 ‘key’(例如,整数或字符串)标识,并持有一个 ‘value’(通常是 NDArray)。
  • 在神经网络训练中,每个参数数组(例如,某层的权重)都被分配一个 key。其 value 是包含参数数据的实际 NDArray。
  • 工作节点(执行计算的设备)在处理完一批数据后向 KVStore ‘push’(推送)梯度。然后在处理下一批数据之前从 KVStore ‘pull’(拉取)更新后的参数。

本质上,KVStore 充当一个集中式或分布式仓库,使得各种计算单元之间能够高效地共享和同步数据。

将 KVStore 想象成一个可供不同设备访问的共享内存对象。每个设备可以将数据 ‘push’(写入)到其中,并从中 ‘pull’(读取)数据。

我们来逐步了解基本的实现步骤:

初始化:首先,我们在 KVStore 中初始化 key-value 对。这里,我们将初始化一个 key(例如,整数 3),并为其关联一个 NDArray,然后将这个值拉取出来。

import mxnet as mx
from mxnet import nd
# Create a local KVStore (for single-machine, multi-device scenarios)
# 创建一个本地 KVStore(用于单机、多设备场景)
kv = mx.kv.create('local')
shape = (3, 3)
key = 3
# Initialize key '3' with a 3x3 NDArray of 2s
# 初始化 key '3',使用一个由 2 填充的 3x3 NDArray
kv.init(key, nd.ones(shape) * 2)
# Create an empty NDArray to pull the value into
# 创建一个空的 NDArray,用于将值拉入其中
a = nd.zeros(shape)
kv.pull(key, out=a)
print("Initialized value:")
print(a.asnumpy())

输出:

Initialized value:
[[2. 2. 2.]
[2. 2. 2.]
[2. 2. 2.]]

Push, Aggregate, and Update:初始化后,我们可以向一个 key 推送新值。KVStore 可以聚合多个推送的值,或根据定义的更新器函数更新现有值。

推送一个新值(默认行为是覆盖):

kv.push(key, nd.ones(shape) * 8)
kv.pull(key, out=a) # Pull the updated value
# 拉取更新后的值
print("Value after push:")
print(a.asnumpy())

输出:

Value after push:
[[8. 8. 8.]
[8. 8. 8.]
[8. 8. 8.]]

用于推送的数据可以位于任何设备(CPU 或 GPU)上。如果同时有多个值从不同源(例如,来自多个 GPU 的梯度)推送到同一个 key,KVStore 通常会在应用更新规则之前将这些值相加。我们来模拟从多个上下文(context)推送数据:

# Initialize key '3' again for this example, perhaps to zero or a known state.
# 再次初始化 key '3' 用于本例,可能初始化为零或已知状态。
kv.init(key, nd.zeros(shape)) # Re-initialize to see aggregation clearly
# 重新初始化为零,以便清晰地看到聚合效果
num_devices = 4
# contexts = [mx.gpu(i) for i in range(num_devices)] # If GPUs are available
# 如果 GPU 可用
contexts = [mx.cpu(i) for i in range(num_devices)] # Using CPUs for broad compatibility
# 使用 CPU 以获得更广泛的兼容性
# Each device has a piece of data to push (e.g., gradients)
# 每个设备都有一部分数据要推送(例如,梯度)
# For this example, let's assume each pushes an array of ones.
# 对于本例,假设每个设备都推送一个全一的数组。
# The KVStore will sum these by default if pushed in a list for a single key.
# 如果将它们作为一个列表推送给单个 key,KVStore 默认会将它们相加。
arrays_to_push = [nd.ones(shape, ctx=contexts[i]) for i in range(num_devices)]
kv.push(key, arrays_to_push)
kv.pull(key, out=a)
print("Value after pushing from multiple sources (summed):")
print(a.asnumpy())

输出:

Value after pushing from multiple sources (summed):
[[4. 4. 4.]
[4. 4. 4.]
[4. 4. 4.]]

Custom Updater:默认情况下,推送一个单独的 NDArray 会覆盖现有值(或如果推送的是列表,则使用预定义的聚合策略,如求和)。我们可以定义一个自定义更新器函数来控制新值如何修改现有值。更新器在 push 操作期间被调用。

# Define a custom updater
# 定义一个自定义更新器
def custom_updater(key_id, input_value, stored_value):
print(f"Updater called on key: {key_id}")
# Example update rule: stored_value = stored_value + input_value * 2
# 示例更新规则:stored_value = stored_value + input_value * 2
stored_value[:] += input_value * 2 # Update in-place
# 原地更新
kv.set_updater(custom_updater)
# Current value in 'a' and KVStore for key '3' is [[4., 4., 4.]]
# 'a' 中和 KVStore 中 key '3' 的当前值是 [[4., 4., 4.]]
# Let's push a new value
# 让我们推送一个新值
kv.push(key, nd.ones(shape))
# The updater will be called: stored_value (4) += input_value (1) * 2 => 4 + 2 = 6
# 更新器将被调用:存储值 (4) += 输入值 (1) * 2 => 4 + 2 = 6
kv.pull(key, out=a)
print("Value after push with custom updater:")
print(a.asnumpy())

输出(来自更新器的打印可能会出现在最终打印之前):

Updater called on key: 3
Value after push with custom updater:
[[6. 6. 6.]
[6. 6. 6.]
[6. 6. 6.]]

Pull:与 push 类似,可以通过一次调用将值拉取到多个设备上。out 参数可以是一个 NDArray 的列表,每个 NDArray 位于不同的上下文(context)上。

b_list = [nd.zeros(shape, ctx=contexts[i]) for i in range(num_devices)]
kv.pull(key, out=b_list)
print("Value pulled onto one of the devices:")
print(b_list[1].asnumpy()) # Displaying the value on the second device/context
# 显示第二个设备/上下文上的值

输出:

Value pulled onto one of the devices:
[[6. 6. 6.]
[6. 6. 6.]
[6. 6. 6.]]
import mxnet as mx
from mxnet import nd
kv = mx.kv.create('local')
shape = (3, 3)
key = 3 # Using a single key for simplicity in this example
# 在本例中为了简单使用单个 key
# 1. Initialization
# 1. 初始化
print("1. Initialization")
kv.init(key, nd.ones(shape) * 2)
a = nd.zeros(shape)
kv.pull(key, out=a)
print(a.asnumpy())
# 2. Push (overwrite)
# 2. 推送(覆盖)
print("\n2. Push (overwrite)")
kv.push(key, nd.ones(shape) * 8)
kv.pull(key, out=a)
print(a.asnumpy())
# 3. Push from multiple sources (aggregation - sum by default)
# 3. 从多个源推送(聚合 - 默认为求和)
print("\n3. Push from multiple sources (aggregation)")
# Re-initialize for clarity of aggregation, or ensure previous value is known.
# 重新初始化以清晰显示聚合,或确保已知之前的值。
kv.init(key, nd.zeros(shape)) # Initialize to zeros
# 初始化为零
num_devices = 4
contexts = [mx.cpu(i) for i in range(num_devices)]
arrays_to_push = [nd.ones(shape, ctx=c) for c in contexts]
kv.push(key, arrays_to_push)
kv.pull(key, out=a)
print(a.asnumpy())
# 4. Custom Updater
# 4. 自定义更新器
print("\n4. Custom Updater")
def custom_updater_example(key_id, input_value, stored_value):
print(f"Updater logic: key {key_id}, input shape {input_value.shape}, stored shape {stored_value.shape}")
stored_value[:] += input_value * 2 # Modify stored_value in-place
# 原地修改 stored_value
kv.set_updater(custom_updater_example)
# 'a' currently holds [[4., 4., 4.]] from the previous step.
# 'a' 当前持有 [[4., 4., 4.]],来自上一步。
kv.push(key, nd.ones(shape)) # input_value is ones(shape)
# input_value 是 ones(shape)
# Updater: stored_value (4) becomes 4 + (1*2) = 6
# 更新器:stored_value (4) 变成 4 + (1*2) = 6
kv.pull(key, out=a)
print(a.asnumpy())
# 5. Pull to multiple devices
# 5. 拉取到多个设备
print("\n5. Pull to multiple devices")
b_list = [nd.zeros(shape, ctx=c) for c in contexts]
kv.pull(key, out=b_list)
print(b_list[1].asnumpy()) # Value on one of the devices
# 设备之一上的值

KVStore 高效地处理对 key-value 对列表的操作。

在单个设备上下文(device context)上初始化、推送和拉取多个 key 的示例:

keys = [5, 7, 9]
shape = (2,2) # Using a different shape for variety
# 使用不同的形状以示多样性
# Initialize multiple keys
# 初始化多个 keys
kv.init(keys, [nd.ones(shape) for _ in keys])
# Push new values to these keys (using the custom_updater from before if still set)
# 将新值推送到这些 keys(如果之前设置了 custom_updater,将继续使用)
# To avoid unintended effects from previous custom_updater, let's use default updater for this section
# 为避免之前 custom_updater 产生的意外影响,让我们在本节使用默认更新器
# For default ASSIGN behavior with individual pushes:
# 对于使用单独推送的默认 ASSIGN 行为:
# kv.set_updater(None) # Or, more explicitly, mx.optimizer.Optimizer.get_updater(mx.optimizer.create('sgd')) would yield a no-op essentially for basic push/pull if no training.
# 或者,更明确地说,mx.optimizer.Optimizer.get_updater(mx.optimizer.create('sgd')) 本质上会为基本 push/pull 操作生成一个空操作(如果不是训练)。
# For simple overwrite, push new values. If custom updater still active, it will apply.
# 对于简单的覆盖,只需推送新值即可。如果自定义更新器仍处于活动状态,它将生效。
# Let's assume default behavior or re-init KVStore for clean state.
# 对于本例,让我们假设是全新的 KV 或一个更新器仅执行加法操作的 KV。
# If using the custom_updater (stored += input * 2):
# 如果使用 custom_updater (stored += input * 2):
# init_val = 1. pushed_val = 1. result = 1 + 1*2 = 3.
# 初始值 = 1,推送值 = 1。结果 = 1 + 1*2 = 3。
kv.push(keys, [nd.ones(shape) for _ in keys])
outputs_list = [nd.zeros(shape) for _ in keys]
kv.pull(keys, out=outputs_list)
print("\nValue for one of the keys (e.g., key 7) after multi-key operations:")
# The index in outputs_list corresponds to the index in keys list
# outputs_list 中的索引对应于 keys 列表中的索引
print(outputs_list[1].asnumpy()) # Corresponds to key 7
# 对应于 key 7

输出(假设 custom_updater stored_val[:] += input_val * 2 处于活动状态,且初始化值为 1,推送值为 1):

Updater logic: key 5, input shape (2,2), stored shape (2,2)
Updater logic: key 7, input shape (2,2), stored shape (2,2)
Updater logic: key 9, input shape (2,2), stored shape (2,2)
Value for one of the keys (e.g., key 7) after multi-key operations:
[[3. 3.]
[3. 3.]]

推送多个 key 的数据,其中每个 key 的数据可能被分片(sharded)或来自多个设备:

keys = [10, 11, 12] # New set of keys
# 新的一组 keys
shape = (2,2)
num_devices = 2
contexts = [mx.cpu(i) for i in range(num_devices)]
# Initialize keys
# 初始化 keys
kv.init(keys, [nd.ones(shape) * 0.5 for _ in keys]) # Initial value 0.5
# 初始值 0.5
# Data to push: for each key, a list of NDArrays, one from each device
# 要推送的数据:对于每个 key,是一个 NDArray 列表,每个列表项来自一个设备
# E.g., key 10 gets [data_dev0_key10, data_dev1_key10]
# 例如,key 10 接收 [data_dev0_key10, data_dev1_key10]
# If custom_updater (val += input*2) is active, and sum aggregation occurs first:
# 如果 custom_updater (val += input*2) 处于活动状态,并且首先发生输入相加:
# For key 10: dev0 pushes 1, dev1 pushes 1. Aggregated input = 2.
# 对于 key 10:设备 0 推送 1,设备 1 推送 1。聚合后的输入 = 2。
# Stored = 0.5. New_stored = 0.5 + 2*2 = 4.5
# 存储的值 = 0.5。新的存储值 = 0.5 + 2*2 = 4.5
data_to_push_multi_device = []
for _ in keys:
key_data_from_devices = [nd.ones(shape, ctx=c) for c in contexts]
data_to_push_multi_device.append(key_data_from_devices)
kv.push(keys, data_to_push_multi_device)
# Pull to multiple devices (or a list of NDArrays on one device)
# 拉取到多个设备(或一个设备上的 NDArray 列表)
pulled_data_multi_device = [[nd.zeros(shape, ctx=c) for c in contexts] for _ in keys]
kv.pull(keys, out=pulled_data_multi_device)
print("\nValue for one key (e.g., key 11) on one device (e.g., device 1) after multi-device push:")
print(pulled_data_multi_device[1][1].asnumpy()) # Data for key 11, on context 1
# key 11 的数据,在上下文 1 上

输出(假设 custom_updater stored_val[:] += input_val * 2 处于活动状态,并且对输入进行求和聚合):

Updater logic: key 10, input shape (2,2), stored shape (2,2)
Updater logic: key 11, input shape (2,2), stored shape (2,2)
Updater logic: key 12, input shape (2,2), stored shape (2,2)
Value for one key (e.g., key 11) on one device (e.g., device 1) after multi-device push:
[[4.5 4.5]
[4.5 4.5]]

MXNet 的可视化工具,例如 mxnet.viz.plot_network,通过将神经网络渲染为计算图来帮助理解其架构。该图将层或操作显示为节点(nodes),将数据流显示为边(edges)。虽然这通常需要安装 Graphviz 库并生成图像,但也可以通过检查网络的符号定义来理解底层结构。

我们可以使用 mx.viz.plot_network 来可视化符号定义(symbolically defined)的网络。先决条件:

  • 一个 Python 环境(例如,Jupyter Notebook 或脚本)。
  • Graphviz 库(包括核心库和 Python graphviz 包)。您通常可以通过 pip install graphviz 和系统包管理器(例如,在 Debian/Ubuntu 上使用 apt-get install graphviz)安装它。

我们来定义一个用于线性矩阵分解的简单网络,并描述其结构,就像被可视化一样。

import mxnet as mx
from mxnet import sym
# Define symbolic variables for inputs
# 定义输入的符号变量
user = sym.Variable('user')
item = sym.Variable('item')
score = sym.Variable('score') # Label for training
# 训练用的标签
# Hyperparameters (dummy dimensions for example)
# 超参数(示例用的虚拟维度)
k = 64 # Dimensionality of latent factors
# 潜在因子的维度
max_user_id = 1000 # Number of unique users
# 唯一用户数量
max_item_id = 500 # Number of unique items
# 唯一物品数量
# User embedding layer
# 用户嵌入层
user_embedding = sym.Embedding(data=user, input_dim=max_user_id, output_dim=k, name='user_embedding')
# Item embedding layer
# 物品嵌入层
item_embedding = sym.Embedding(data=item, input_dim=max_item_id, output_dim=k, name='item_embedding')
# Predict score via inner product of user and item embeddings
# 通过用户和物品嵌入的内积预测得分
pred = user_embedding * item_embedding
pred = sym.sum_axis(data=pred, axis=1, name='sum_ reducción') # Corrected name for clarity
# 为清晰起见修正名称
pred = sym.Flatten(data=pred, name='flatten_output')
# Define the loss function (e.g., Linear Regression Output for mean squared error)
# 定义损失函数(例如,用于均方误差的线性回归输出)
# This also makes 'score' the label for this network output
# 这也将 'score' 设置为该网络输出的标签
network_symbol = sym.LinearRegressionOutput(data=pred, label=score, name='lro')
# To actually plot (if Graphviz is installed):
# 实际绘图(如果安装了 Graphviz):
# viz_graph = mx.viz.plot_network(network_symbol)
# viz_graph.render("matrix_factorization_network", view=False) # Saves to PDF
# 保存为 PDF
print("Textual description of the network symbol:")
# 网络符号的文本描述:
print("Inputs: 'user', 'item' (features); 'score' (label).")
# 输入:'user', 'item'(特征);'score'(标签)。
print(f"1. 'user' -> Embedding (name='user_embedding', input_dim={max_user_id}, output_dim={k})")
# 1. 'user' -> 嵌入层 (名称='user_embedding', 输入维度=..., 输出维度=...)
print(f"2. 'item' -> Embedding (name='item_embedding', input_dim={max_item_id}, output_dim={k})")
# 2. 'item' -> 嵌入层 (名称='item_embedding', 输入维度=..., 输出维度=...)
print("3. Outputs of user_embedding and item_embedding -> Element-wise Multiplication")
# 3. user_embedding 和 item_embedding 的输出 -> 元素级乘法
print("4. Result -> Sum along axis 1 (name='sum_reducción')")
# 4. 结果 -> 沿轴 1 求和 (名称='sum_reducción')
print("5. Result -> Flatten (name='flatten_output')")
# 5. 结果 -> 展平 (名称='flatten_output')
print("6. Result (prediction), 'score' (label) -> LinearRegressionOutput (name='lro') for loss calculation.")
# 6. 结果(预测)、'score'(标签)-> LinearRegressionOutput (名称='lro') 用于计算损失。
print("This network structure represents a common matrix factorization model.")
# 这个网络结构代表了一个常见的矩阵分解模型。
# For more detailed information on specific layers and their connections,
# 对于特定层及其连接的更详细信息,
# one could iterate through network_symbol.get_internals() or analyze the JSON representation:
# 可以遍历 network_symbol.get_internals() 或分析 JSON 表示形式:
# print(network_symbol.tojson())
# print(network_symbol.tojson())

运行 mx.viz.plot_network 将生成一个图,显示这些变量和操作作为节点,箭头指示数据流从输入通过嵌入、乘法、求和、展平,最终到达损失输出层。这种视觉表示对于调试和理解复杂的模型架构非常宝贵。