SPSALearningRateCalibrationConfig¶
from cascaqit import SPSALearningRateCalibrationConfig
SPSALearningRateCalibrationConfig ¶
SPSALearningRateCalibrationConfig(
directions: int = 5,
target_update_rms: float = 0.1,
gradient_rms_floor: float = 1e-08,
max_learning_rate: float | None = None,
schema_version: str = ALGORITHM_SCHEMA_VERSION,
)
在 SPSA 正式迭代前估计梯度尺度,据此选择学习率。directions 是校准方向数,必须为正整数;target_update_rms 是希望首步达到的更新均方根,必须为正数。校准会实际调用目标函数,消耗求值与后端预算。
对校准得到的平均梯度,记各参数分量的均方根为 g。学习率尺度为 target_update_rms * (1 + stability_constant)**learning_rate_exponent / g。g 低于正的 gradient_rms_floor 时会报错;这不是把分母截到下限。若计算出的学习率超过可选正数 max_learning_rate,也会报错,不会自动裁剪。
将它放入 SPSAConfig.learning_rate_calibration,同时设 learning_rate=None。边界投影与估计噪声可能使实际更新不同于目标均方根,应查看保存的校准记录。
from cascaqit import SPSAConfig, SPSALearningRateCalibrationConfig
calibration = SPSALearningRateCalibrationConfig(directions=3, target_update_rms=0.05)
config = SPSAConfig(learning_rate=None, learning_rate_calibration=calibration)
assert config.learning_rate is None
assert SPSALearningRateCalibrationConfig.from_json(calibration.to_json()) == calibration
从字典还原 SPSALearningRateCalibrationConfig。省略字段使用默认值,嵌套配置还原后重新校验。 缺少必需字段或字段不合法时可能抛出 KeyError、TypeError 或 ValueError。
解析 JSON 对象并调用 from_dict(),返回 SPSALearningRateCalibrationConfig。非法 JSON 会抛出解析错误;顶层不是对象时抛出 TypeError。
返回可写入 JSON 的字典,嵌套对象一并序列化。元组转成数组;这个字典是保存的声明,不是执行结果。
返回 JSON 字符串,不写文件。indent=None 使用紧凑格式;提供缩进宽度可便于阅读。
返回规范 JSON 的 SHA-256 十六进制摘要。字段、标识或元数据变化都可能改变摘要;它用于比较保存内容,不判断两个声明在物理上是否等价。