
人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载本指南以google-research仓库中 incontext/README.md 所指向的预印本研究《What Learning Algorithm is In-context Learning? Investigations with Linear Models》为骨架系统拆解 incontext/ 代码库的全部实现数据生成、Transformer 训练、对比算法、探针模型与可视化工具。读者读完本文后将掌握这套实验框架的完整调用链、全部命令行参数含义并能复现上下文内学习In-context LearningICL究竟等价于哪种经典学习算法这一核心实验。一、研究背景与代码库定位大语言模型能够在提示中通过少量示例few-shot exemplars完成新任务而不更新权重这一现象被称为 In-context Learning。一个核心科学问题是Transformer 在上下文内执行的隐式学习算法到底是什么incontext/README.md 明确指出本仓库对应预印本论文What Learning Algorithm is In-context Learning? Investigations with Linear ModelsEkin Akyürek, Jacob Andreas, Dale Schuurmans, Tengyu Ma, Denny Zhou该研究采用用线性模型做实验的策略把 ICL 问题简化为线性回归的上下文内学习——即给定若干(x, y)示例后让 Transformer 预测新输入x对应的y。由于线性回归存在解析解最小二乘、岭回归等研究者可以把 Transformer 在上下文内隐式学习的行为权重behavioral weight与经典算法逐一对照从而判定其内部实现的算法类型。代码库本身是一个可运行的实验框架由以下模块构成目录结构见 incontext/文件职责incontext/incontext/sampler_lib.py线性回归数据的随机采样与分布字符串解析incontext/incontext/transformer_lib_flax.pyFlax Transformer 主干与学习率调度器incontext/incontext/predictor_flax.py因果语言模型CausalLM预测头与损失计算incontext/incontext/model_trainer.py主模型训练、评估与可视化入口incontext/incontext/main.py实验入口参数解析、训练调度、探针可选训练incontext/incontext/algos.py经典回归算法集合与在线回归评测incontext/incontext/probe_flax.py / probe_trainer.py从隐状态解码权重的探针模型incontext/incontext/plotting.py经验分布、隐式权重路径等可视化incontext/incontext/linear_model.py / train_linear_lib.py单层线性模型的 SGD/GD 基准训练incontext/incontext/utils.py随机种子、Flag 转配置等公共工具incontext/run.sh / requirements.txt / setup.py环境搭建与打包依赖以 JAX 生态为主requirements.txt 声明了flax、jax、matplotlib、tensorflow用于gfile文件读写、numpy、absl-py、optax此外 algos.py 还直接使用了sklearn的linear_model、neighbors、metrics模块。二、数据生成Sampler 与分布字符串语法2.1 分布字符串的解析规则训练与评估的数据分布完全由字符串控制解析逻辑在 sampler_lib.py 的str_to_distribution_fn中单分布格式为类型*alphabeta类型仅支持uniformnp.random.rand与normalnp.random.randn采样结果乘以alpha并加上beta。例如默认值normal*1.00.0即标准正态分布normal*2.50表示标准差放大到 2.5normal*12.5表示均值偏移到 2.5。混合分布用逗号连接多个单分布例如normal*10,uniform*10。此时每个子分布各采样一份 batch再按 batch 维度随机挑选合并等效于均匀混合采样。未知类型会抛出ValueError(Unknown distribution type.)。2.2 Sampler 的序列构造Sampler类sampler_lib.py负责生成一条完整的上下文序列核心参数包括length示例exemplar个数dimx向量的维度hidden_size序列向量的维度填充后的宽度x_distribution_fn/w_distribution_fnx与回归系数w的采样函数noise_std大于 0 时给y添加高斯噪声默认 0.0。sample()方法sampler_lib.py的核心逻辑是为每个示例依次采样x与系数w通过np.einsum(bi,bi-b, coefficients, x)计算y然后将x_vec[0, x]前置零位与y_vec[y, 0...0]后置dim个零填充交替拼接成一条长度为length * 2的序列返回(out, coefficients, xs, ys)四元组。这种x 与 y 交替、零填充对齐维度的编码方式与后续CausalLM中按偶数位抽取y的损失设计一一对应详见第三节。2.3 对训练与评估的影响noise_std直接控制任务难度在 model_trainer.py 中主训练用的Sampler以x_distribution_str、w_distribution_str、noise_std构造而在评估阶段 model_trainer.py 会额外以eval_noise_std x_dim / 20生成带噪声的测试分布用于考察模型对标签噪声的鲁棒性。此外 test_empirical_statistics 中采样num_exemplars - 1个示例并注释说明否则用于最后一个示例的位置编码将是未训练参数——这是序列长度与位置编码最大长度max_len(num_exemplars1)*2之间必须留出余量的直接证据。三、Transformer 实现CausalLM 与自注意力补丁3.1 CausalLM按偶数位取 y 的自回归损失incontext/incontext/predictor_flax.py 定义了CausalLM模块。前向流程为取seq_from inputs[:, :-1, :]预测下一位置用nn.attention.make_causal_mask构造因果掩码送入 transformer_lib_flax.py 的Transformer可返回各层注意力权重经nn.Dense投影默认输出维度 1即只预测 y开启loss_on_x_steps时预测完整向量extract_ypredictor_flax.py按jnp.arange(offset, seq.shape[1], 2)取出偶数位置即只保留 y 步的预测损失为预测 y 与目标 y 的平方误差和(y_pred - y_target)**2).sum(axis-1)可选的return_attentionTrue会把各层注意力权重一并返回供可视化使用。loss_on_x_steps标志定义在 predictor_flax.py默认False表示只对 y 步计算损失。3.2 Transformer 主干与可配置项transformer_lib_flax.py 的TransformerConfig定义了全部结构超参数而对应的命令行 Flag 定义在文件开头L45-L73Flag默认值含义n_layers12编码器层数n_heads8注意力头数要求hidden_size % n_heads 0hidden_size512模型宽度norm_firstTruePre-LN先 LayerNorm 再注意力/MLPfinal_layer_normFalse是否在输出前追加最后一层 LayerNormdisable_layer_normsFalse完全禁用 LayerNorm消融用inner_dimNoneMLP 中间维度None 时取hidden_size * 4kernel_init/bias_inituniform_scaling注意力/MLP 参数初始化器linear_w_init/linear_bias_inituniform_scaling投影层初始化器posemb_inituniform_scaling位置嵌入初始化器activation_fngelu激活函数可选 gelu/relu/tanh/softplus/sofuTransformer的组成L355-L408先经nn.Dense投影到hidden_size再加可学习位置嵌入PositionEmbeddingsL189-L221按max_len截取然后堆叠num_layers个Encoder1DBlockL267-L352每个块为自注意力 残差 MLP 残差结构。初始化器支持字符串解析nn_init_parserL127-L140uniform_scaling、ones、zeros、normal(std)。3.3 自注意力补丁返回注意力权重incontext/flax/self_attention_patch.py 是从flax.linen.attention复制并改造的版本在 dot_product_attention 与MultiHeadDotProductAttention中增加了return_attention分支使注意力权重可以被返回并序列化。Encoder1DBlock调用的是补丁版SelfAttentiontransformer_lib_flax.py其qkv_featureshidden_size // num_heads。这一改动直接服务于 model_trainer.py 中导出attention_stats.pkl的分析需求——研究者借此观察注意力在示例之间如何分配权重。四、训练管线与命令行参数全解4.1 实验入口 main.pyincontext/incontext/main.py 是标准实验入口流程为解析 flags →utils.flags_to_args()生成ConfigDict→ 创建实验目录并写config.json→utils.set_seed→ 初始化模型 → 训练主模型 → 保存metrics.pickle与 checkpoint → 可选训练探针。其自有 FlagFlag默认值含义seed0随机种子同时作用于 Python、numpy 与 JAXbatch_size64每步采样的序列条数x_dim20输入维度也是系数维度num_exemplars40每条序列的示例个数exp_folderexp实验输出目录debugFalse调试预测与后验分布train_probeFalse主模型训练完成后是否追加训练探针主模型训练与评估 Flag 定义在 model_trainer.pyFlag默认值含义n_epochs5001训练轮数n_iter_per_epoch100每轮迭代次数learning_rate1e-4基础学习率weight_decay0AdamW 权重衰减lr_scheduler_typecosinecosine / warmup / 常数adam_b10.9Adam 一阶动量adam_b20.98Adam 二阶动量adam_eps1e-9Adam 数值稳定项x_distribution_strnormal*1.00.0训练时 x 的分布w_distribution_strnormal*1.00.0训练时 w 的分布noise_std0.0训练标签噪声标准差4.2 训练步骤与多设备并行train_stepmodel_trainer.py实现单步更新jax.value_and_grad求梯度jax.lax.pmean(grads, batch)做跨设备梯度平均state.apply_gradients更新参数损失经pmean汇总、y_errors经psum求和。模型通过jax.pmap并行化L317-L320数据用common_utils.shard分片到各设备dropout 随机数按设备拆分L418-L420。优化器为optax.adamwL301-L307学习率调度有三种选择由lr_scheduler_type控制L284-L299cosinecreate_learning_rate_schedulertransformer_lib_flax.pywarmup 步数为(n_epochs // 5) * n_iter_per_epoch随后按半余弦周期衰减warmupcreate_learning_rate_scheduler_v2L448-L516因子串constant * linear_warmup其他恒为learning_rate的常数调度。train主循环model_trainer.py每eval_every_n_epochs默认 1000触发一次eval_model并把每轮指标堆叠保存checkpoint 经flax.training.checkpoints.save_checkpoint写入exp_folder/ckpt/L226-L248。4.3 跨分布评估eval_modelmodel_trainer.py在固定测试分布集合上分别扰动w 分布与x 分布进行评估。默认测试集为(normal*10, normal*2.50, normal*12.5)当训练分布本身是混合分布x_distribution_str含逗号时还会追加normal*1-2.5与normal*15.0L431-L444。每轮评估输出到exp_folder/plots/distribution/w_{str}/与x_{str}/子目录其中*被替换为x以兼容文件系统x 扰动下还会生成noise_{x_dim/20}/子目录对应带噪测试。五、对比算法集合ICL 的候选学习算法5.1 统一接口与评分incontext/incontext/algos.py 定义了RegressionAlgorithm抽象基类L50-L120统一接口为fit(x, y)、predict(x)、get_parameters()、iterate(x, y)、is_iterative()、reset()并提供scores()计算 R2 与 MSE。实现包括LeastSquareAlgorithmL123-L148sklearn.linear_model.LinearRegression即标准最小二乘FakeLeastSquareAlgorithmL151-L180直接使用先验精度矩阵计算weight (precision x.T y) / n的伪最小二乘对应文中Lstsq-Constant-SigmaRidgeRegressionAlgorithmL183-L214sklearn.linear_model.Ridgealpha默认 0.01KNNAlgorithmL217-L243sklearn.neighbors.KNeighborsRegressor支持 uniform / distance 加权SGDL246-L306自实现随机梯度下降window1表示在线只看当前样本window-1表示全量梯度梯度裁剪到[-20, 20]支持weight_decay。5.2 在线回归评测核心评测函数online_regressionalgos.py模拟上下文内学习过程对每条序列遍历第 1 到最后一个示例迭代式算法is_iterative()用当前样本iterate非迭代算法用fit(x[:i])重拟合随后预测下一个x_i并记录预测与参数。online_regression_with_batchL332-L349批量执行并返回逐示例的 MSE 曲线。值得一提的细节algos.py的if __name__ __main__分支L352-L503)是一个可直接运行的算法对比实验——在x_dim10, hidden_size128, num_exemplars64下对 16 种算法变体不同学习率、窗口、权重衰减、Ridge alpha、KNN 加权在三种 w/x 分布下绘制 MSE 曲线保存为algos_w_*.jpeg/algos_x_*.jpeg。这是先有候选算法、再与 Transformer 行为对齐方法论的直接体现。六、探针模型从隐状态解码真实权重除了行为层面的对比代码库还提供了权重层面的探针直接训练一个线性解码器从 Transformer 各层最后一个 token 的隐状态恢复回归系数。incontext/incontext/probe_flax.py 的ProbeModel对每一层独立使用一个nn.Dense(config.x_dim)输入该层最后一个位置的 hidden state输出与真实系数coefficients对比求平方误差最终返回形状为 batch 的逐层误差堆叠。ProbeConfigL42-L51包含hidden_size、num_layers、x_dim、max_len等字段。incontext/incontext/probe_trainer.py 负责探针训练用冻结的主模型state.params前向得到seq_hiddensL115-L117)转置为(layer, batch, len, hidden)后分片训练循环与主模型一致采用optax.adamw与 cosine 调度L189-L212新增 Flagprobe_learning_rate默认 0.001、probe_epochs默认 20、probe_iters默认 100、probe_lr_scheduler_type默认 cosine。探针的意义在于如果某层隐状态能被线性映射出精确的回归系数就说明该层内部编码了最小二乘等算法对应的解为Transformer 隐式执行了某类学习算法提供内部证据。七、可视化与实证分析工具incontext/incontext/plotting.py 提供三类核心可视化全部由 model_trainer.py 的test_empirical_statistics触发输出在exp_folder/plots/下行为权重的经验拟合plot_empirical_distributionL75-L322用模型对一组随机x的预测y_pred反拟合一个带截距的最小二乘模型algo_fn(fit_interceptTrue)得到行为权重w_empirical与 R2随后绘制3D 经验平面红色与真实平面蓝色对比plot_planes_3dL62-L72逐示例的预测误差曲线_errors.jpeg、与最小二乘/岭回归/SGD 等算法的误差对比plot_an_algorithm调用online_regression_with_batch各算法与 Transformer 预测的发散度曲线_divergence.jpeg可选的预测序列保存_predictions.pkl。隐式权重路径plot_implicit_wL325-L382在每个示例前缀长度上重复预测→反拟合过程绘制w_empirical在权重空间中的轨迹_w.jpeg、各前缀 R2_r2.jpeg以及与真实最小二乘权重差异||Wlsq2-Wimp||^2随示例数变化的曲线_wdiff.jpeg——这是判定ICL 是否收敛到最小二乘最直观的证据图。基向量可视化plot_basis_imageL385-L428逐个维度构造只有该维度非零的x比较模型预测红点与真实y蓝点的散点图检查模型是否在每个坐标方向上对齐。此外average_stats.jpeg、average_stats.pkl与attention_stats.pkl汇总了平均预测误差与各层注意力权重供后续离线分析。八、线性基准单层模型的 SGD/GD 对照为回答ICL 等价于什么算法incontext/incontext/linear_model.py 与 incontext/incontext/train_linear_lib.py 还提供了一个显式线性层 SGD/GD的对照组LinearModellinear_model.py只是一个无偏置的nn.Dense(1)LinearConfig.alpha控制 L2 正则强度。train_linear_lib.py 的run_expL98-L163在每条序列上运行optax.sgd其中gdTrue时用全量历史样本计算梯度梯度下降gdFalse时只用当前样本真正的 SGD并逐示例记录下一示例的预测损失。mainL166-L202默认在x_dim10、num_exemplars32、w 分布normal*15.0下采样并平均多条序列的损失输出losses.png。该实验用于刻画显式 SGD 在同样任务上的学习曲线作为与 Transformer 隐式行为对照的基准。九、环境搭建与运行指南9.1 依赖安装仓库提供两种方式一键脚本run.sh 执行virtualenv -p python3 .创建本地虚拟环境、source ./bin/activate激活、pip install -r requirements.txt安装依赖随后运行python -m incontext.utils验证模块可导入。注意该脚本基于 Bash 与virtualenv需要预先安装对应工具。手动安装直接pip install -r requirements.txt或执行pip install .setup.py 以setuptools将包注册为incontext。值得提醒requirements.txt未显式列出scikit-learn但 algos.py 与 least_square.py 均直接import sklearn因此运行对比算法实验前需自行补装scikit-learn。模型训练需要 JAX 可用的后端CPU/GPU/TPU 均可多设备时自动启用pmap并行。9.2 复现 Transformer 主实验以 main.py 的默认参数x_dim20、num_exemplars40、batch_size64、n_epochs5001、hidden_size512、n_layers12、n_heads8为例# 训练 评估含绘图与 checkpoint python -m incontext.main \ --seed0 \ --x_dim20 \ --num_exemplars40 \ --batch_size64 \ --n_epochs5001 \ --exp_folderexp # 附加探针训练 python -m incontext.main --train_probeTrue --exp_folderexp # 修改训练分布如增大 w 方差 python -m incontext.main --w_distribution_strnormal*2.50 --exp_folderexp_w2p5运行产物集中在exp_folderconfig.json参数快照、metrics.pickle逐步指标、ckpt/可恢复的 checkpoint、plots/distribution/全部可视化与统计数据。9.3 运行算法对比与线性基准# 经典算法集合的 MSE 曲线直接执行 algos.py 的 main 分支 python -m incontext.algos # 单层线性模型 SGD/GD 基准 python -m incontext.train_linear_lib --seed0 --x_dim10 --num_exemplars32 --batch_size32 # 最小二乘变体对比示例least_square.py 的 main 分支含伪最小二乘 python -m incontext.least_square十、总结这套框架回答了什么问题纵观全库incontext/ 的贡献是一套**算法归因实验方法论**在可控的线性回归 ICL 任务上同时实现a数据分布参数化采样第二节、b可配置的因果 Transformer 训练与跨分布评估第三、四节、c一组候选经典算法与在线回归评测第五节、d权重层面的探针解码第六节、以及e行为权重路径与误差发散可视化第七节。通过对比 Transformer 的行为权重、隐式权重轨迹与最小二乘/SGD 等算法的对应关系即可实证检验论文提出的核心主张——ICL 内部执行的隐式学习算法究竟与哪一类经典算法一致。需要说明的是本仓库提供的是实验代码与工具链论文的完整实验结论与数据需以预印本正文为准本文所有实现细节均来自 incontext/README.md 及上述源码文件可直接对照查阅。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐PyTorch深度学习中的Transformer语言模型解析PyTorch深度学习中的Transformer语言模型解析 引言语言模型的革命性突破 在深度学习领域Transformer架构的出现彻底改变了自然语言处理示例工程深度强化学习教程Q-Learning算法详解深度强化学习教程Q Learning算法详解 引言 Q Learning是强化学习领域最经典且实用的算法之一它不需要完全了解环境的动态特性即马尔可夫决策过文档教程人工智能深度学习NLP计算机视觉强化学习上一篇CANN/asc-devkit浮点转整型函数下一篇Tutu.ru PHP 后端开发测试题全解析从基础算法到购物车折扣、分布式 Cron 与 Telegram Bot 的完整题解创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考