扩展 - 上三角矩阵演示
1 课程概览
本课讲解上三角矩阵。上三角是对角线上方有值,下方为零。用于掩码机制,1 表示被遮掩(看不见),0 表示可见。生成第 2 个词时,第 2 个词及以后的词都看不见(1),第 1 个词可见(0)。越往后生成,前面看到的词越多。
2 核心概念与定义
- 上三角矩阵:对角线上方有值,下方为零。
- 下三角矩阵:对角线下方有值,上方为零。
- 对角线:从左上到右下的对角线。
- 掩码(mask):1 表示被遮掩(看不见),0 表示可见。
- 时间步:生成字符时,一个时间步一个时间步地解码。
3 模型与算法详解
上三角 vs 下三角
| 类型 | 有值位置 | 零的位置 |
|---|---|---|
| 上三角 | 对角线上方 | 对角线下方 |
| 下三角 | 对角线下方 | 对角线上方 |
记忆方法:哪边有数就叫什么三角。
掩码机制
0 表示能看见,1 表示看不见(被遮掩)。
示例:4 个时间步
生成第 1 个词时:[1, 1, 1, 1] # 所有词都看不见
生成第 2 个词时:[0, 1, 1, 1] # 只看见第 1 个词
生成第 3 个词时:[0, 0, 1, 1] # 看见第 1、2 个词
生成第 4 个词时:[0, 0, 0, 1] # 看见第 1、2、3 个词
生活类比
越往后生成,前面看到的词越多。
- 3 岁/7 岁:人生经历少,很多事看不透
- 30 岁/70 岁:人生经历多,很多事能看透
- 生成第 1 个词:什么都不知道
- 生成最后一个词:前面全知道
应用场景
解码器生成字符时,一个时间步一个时间步地解码。
- 生成第 2 个词时,第 2 个词及以后的词都看不见
- 第 1 个词一定存在(已经生成)
- 类比:今天下午会发生什么不知道,但昨天发生的事已经成定局
4 数学原理与推导
上三角矩阵
$$M_{ij} = \begin{cases} 1, & \text{if } j \geq i \ 0, & \text{if } j < i \end{cases}$$
示例(4×4)
$$M = \begin{bmatrix} 1 & 1 & 1 & 1 \ 0 & 1 & 1 & 1 \ 0 & 0 & 1 & 1 \ 0 & 0 & 0 & 1 \end{bmatrix}$$
5 代码示例
import numpy as np
import torch
def dm01_test_triu(size):
"""
生成上三角矩阵
:param size: 矩阵大小
:return: 上三角矩阵
"""
# 生成上三角矩阵
# np.triu: 上三角矩阵,k=1 表示对角线上移一次
m = np.triu(np.ones((size, size)), k=1)
# 转换类型为 uint8
m = m.astype(np.uint8)
print(f"上三角矩阵 ({size}x{size}):")
print(m)
return m
# 测试
if __name__ == "__main__":
# 生成 5x5 的上三角矩阵
result = dm01_test_triu(5)
print(f"\n形状: {result.shape}")
print(f"类型: {result.dtype}")
输出示例
上三角矩阵 (5x5):
[[0 1 1 1 1]
[0 0 1 1 1]
[0 0 0 1 1]
[0 0 0 0 1]
[0 0 0 0 0]]
代码说明
| 代码 | 说明 |
|---|---|
np.triu(matrix, k=1) | 生成上三角矩阵,k=1 表示对角线上移 |
np.ones((size, size)) | 生成全 1 矩阵 |
astype(np.uint8) | 转换类型为 uint8 |
6 重难点与易错提醒
- ❗重点:上三角是对角线上方有值。
- ❗重点:1 表示被遮掩(看不见),0 表示可见。
- ⚠️易错:不要反着记(零在下叫上三角容易晕)。
- ⚠️易错:
np.triu的 k=1 表示对角线上移一次。 - 💡深入理解:越往后生成,前面看到的词越多。
7 课堂问答精选
Q1:什么是上三角矩阵?
A:对角线上方有值,下方为零的矩阵。记忆方法:哪边有数就叫什么三角。
Q2:掩码中 0 和 1 分别表示什么?
A:0 表示能看见(可见),1 表示看不见(被遮掩)。
Q3:为什么需要上三角矩阵?
A:解码器生成字符时,一个时间步一个时间步地解码。生成第 2 个词时,第 2 个词及以后的词都看不见(1),第 1 个词可见(0)。越往后生成,前面看到的词越多。
Q4:如何用代码生成上三角矩阵?
A:用 np.triu(np.ones((size, size)), k=1),k=1 表示对角线上移一次。
8 本课小结
- 上三角:对角线上方有值。
- 0 表示可见,1 表示被遮掩。
- 生成第 2 个词时,第 2 个词及以后都看不见。
- 越往后生成,前面看到的词越多。
- 代码:
np.triu(np.ones((size, size)), k=1)。
9 延伸思考
- 下三角矩阵如何生成?
- 掩码张量如何可视化?