如何在SGD执行过程中周期性记录物理系统能量值?
问题
我找不到匹配我需求的解决方案。我有一个计算物理系统能量的Python函数,为了优化这个能量,我给它应用了SGD算法。目前我能记录每一步SGD的能量值,但想评估SGD的性能——不需要计算代码总执行时间(我知道怎么弄),而是要了解SGD执行过程中能量随时间的变化,比如每秒记录当前计算出的能量值。
我的代码如下:
def run_one_SGD(max_steps, lr, precis, al_range, optim): al_0 = torch.distributions.Uniform(0, al_range).sample((size_basis,n_pairs)).double().requires_grad_(True) if optim == 'SGD': optimizer = torch.optim.SGD([al_0], lr = lr) elif optim == 'Adam': optimizer = torch.optim.Adam([al_0], lr = lr) en_hist = [1.0] time_hist = [] i = 0 precision = 1 start_time = time.perf_counter() while (i < max_steps and precision > precis): optimizer.zero_grad() S, K, P, H, e, eigs = find_en(al_0) e.backward() optimizer.step() en_hist.append(e.clone().detach()) i += 1 print(f'At {optim} step {i} the energy is {e}') precision = abs((en_hist[-2] - en_hist[-1])/en_hist[-1]) end_time = time.perf_counter() time_hist.insert(i, (end_time - start_time)) return en_hist, time_hist
当前代码能记录每一步SGD的执行时间,但我希望在代码执行过程中周期性(比如每秒)记录能量值e。
解决方案
要实现每秒记录当前能量的需求,你可以在循环中跟踪上一次记录的时间,每次迭代时检查当前时间与上一次记录时间的差值是否达到1秒,若是则记录当前能量和对应时间。
修改后的代码如下:
import time import torch def run_one_SGD(max_steps, lr, precis, al_range, optim): al_0 = torch.distributions.Uniform(0, al_range).sample((size_basis,n_pairs)).double().requires_grad_(True) if optim == 'SGD': optimizer = torch.optim.SGD([al_0], lr = lr) elif optim == 'Adam': optimizer = torch.optim.Adam([al_0], lr = lr) en_hist = [1.0] # 存储按时间间隔记录的能量和对应时间戳 time_based_en = [] time_based_timestamps = [] i = 0 precision = 1 start_time = time.perf_counter() # 记录上一次保存能量的时间点 last_record_time = start_time while (i < max_steps and precision > precis): optimizer.zero_grad() S, K, P, H, e, eigs = find_en(al_0) e.backward() optimizer.step() current_energy = e.clone().detach() en_hist.append(current_energy) i += 1 # 检查是否达到1秒的记录间隔 current_time = time.perf_counter() if current_time - last_record_time >= 1.0: time_based_en.append(current_energy) time_based_timestamps.append(current_time - start_time) last_record_time = current_time print(f'At {current_time - start_time:.2f}s, energy is {current_energy}') print(f'At {optim} step {i} the energy is {current_energy}') precision = abs((en_hist[-2] - en_hist[-1])/en_hist[-1]) return en_hist, time_based_en, time_based_timestamps
关键改动说明:
- 新增
last_record_time变量,跟踪上一次记录能量的时间点 - 每次迭代计算当前时间,判断与上一次记录时间的差值是否≥1秒,满足则记录当前能量和从开始到现在的时间差
- 新增
time_based_en和time_based_timestamps列表,分别存储按时间间隔记录的能量值和对应时间戳,方便后续分析能量随时间的变化趋势 - 保留原有按步骤记录能量的逻辑,同时新增按时间间隔的记录输出
内容的提问来源于stack exchange,提问作者Paolo Recchia
相关产品推荐
相关产品推荐

