BoTorch optimize_acqf函数等式约束参数使用示例求助
解决BoTorch v0.2.3中等式约束的两个常见问题
1. 解决RuntimeError: dot : expected both vectors to have same dtype错误
这个错误的核心是约束张量的 dtype 不匹配:
indices(变量索引)必须是整数类型(LongTensor),用于定位变量;coefficients(系数)和rhs(等式右侧值)必须是浮点类型,要与优化变量的 dtype 保持一致(通常是float32或float64)。
错误示例中,coefficients用了默认的Long类型,与浮点型的优化变量做dot运算时触发了 dtype 冲突。
2. 解决ValueError: No feasible point found错误
这个问题通常有两个原因:
- 变量索引错误:BoTorch采用0索引,如果你要约束的是
x2、x4、x6(从1开始计数的变量),对应的索引应该是[1, 3, 5],而非[2,4,6]; - 初始样本未覆盖可行域:默认的
raw_samples=100可能在严格约束下找不到满足条件的点,需要增大样本数量,或者手动提供可行初始点。
正确的等式约束示例代码
import torch as to from botorch.optim import optimize_acqf from botorch.models import SingleTaskGP from botorch.acquisition import ExpectedImprovement from botorch.fit import fit_gpytorch_model from gpytorch.mlls import ExactMarginalLogLikelihood # 初始化示例模型与采集函数(替换为你的实际模型) train_x = to.rand(10, 20, dtype=to.float32) train_y = to.randn(10, 1, dtype=to.float32) model = SingleTaskGP(train_x, train_y) mll = ExactMarginalLogLikelihood(model.likelihood, model) fit_gpytorch_model(mll) acq_func = ExpectedImprovement(model, best_f=train_y.max()) BATCH_SIZE = 1 # 正确定义等式约束:x2 + x4 + x6 = 1(对应0索引的[1,3,5]) indices = to.tensor([1, 3, 5], dtype=to.long) # 整数索引,必须为long类型 coefficients = to.tensor([1.0, 1.0, 1.0], dtype=to.float32) # 浮点系数 rhs = 1.0 # 浮点型右侧值 new_x, value = optimize_acqf( acq_function=acq_func, bounds=to.tensor([[0.0] * 20, [1.0] * 20], dtype=to.float32), q=BATCH_SIZE, num_restarts=10, raw_samples=500, # 增大初始样本量,提升找到可行点的概率 options={"batch_limit": 5, "maxiter": 200, "nonnegative": True}, equality_constraints=[(indices, coefficients, rhs)], sequential=True, ) # 验证约束是否满足 print("约束验证结果:", to.allclose(new_x[0, indices].sum(), rhs))
额外优化:手动提供初始可行点
如果增大raw_samples仍无法找到可行点,可以手动生成满足约束的初始点并传入:
# 生成10个满足约束的初始点 initial_x = to.rand(10, 20, dtype=to.float32) # 调整指定变量,确保它们的和为1 initial_x[:, indices] = to.rand(10, 3, dtype=to.float32) initial_x[:, indices] = initial_x[:, indices] / initial_x[:, indices].sum(dim=1, keepdim=True) # 在optimize_acqf中加入initial_conditions参数 new_x, value = optimize_acqf( # 其他参数保持不变 initial_conditions=initial_x, )
内容的提问来源于stack exchange,提问作者David Siret Marqués
相关产品推荐
相关产品推荐

