高周波エンジニアのためのAI・機械学習入門(GPU編43) PythonのStable-Baselines-3を使ってLCバンドパスフィルタの素子値を強化学習で最適化する。まずはPPO(Proximal Policy Optimization)を試す。
前回で準備ができたので今回は実際に強化学習をやってみよう。まずはPPO(Proximal Policy Optimization)を試す。 PythonのStable-Baselines-3を使えばいろいろなアルゴリズムが簡単に使える。
コードはこちら。deviceがcpuになっているのはGPU編に反するが、こうしないと余計遅いらしい。PPO以外にもSAC、TD3も試せるようになっている。
"""LC BPFをStable-Baselines3で最適化する共通の実行コード。"""
from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
import bpf
import bpf_score
from bpf_env import LCBPFOptimizationEnv
DEFAULT_TIMESTEPS = {"ppo": 5_000, "sac": 20_000, "td3": 20_000}
def default_targets():
return {
# 帯域内の最大 insertion loss に上限制約
"worst_insertion_loss_dB_in_design_band": {
"upper": 8.0,
"weight": 15.0,
"scale": 1.0,
},
# 一部の周波数だけ良くなるのを防止
"passband_ripple_dB_in_design_band": {
"upper": 1.0,
"weight": 10.0,
"scale": 0.5,
},
# 整合性を維持
"min_return_loss_dB_in_design_band": {
"lower": 10.0,
"weight": 5.0,
"scale": 3.0,
},
# 必要な帯域幅を維持
"measured_fbw": {
"lower": 0.09,
"upper": 0.11,
"weight": 5.0,
"scale": 0.01,
},
}
def build_environment(points=801, reset_noise_scale=0.05):
n = 5
f0 = 2.45
fbw = 0.10
f1, f2 = bpf.band_edges_from_f0_fbw(f0, fbw)
lc_base, q_values, _, _ = bpf.synthesize_lc_bpf(
n=n,
f0=f0,
fbw=fbw,
z0=50.0,
prototype="chebyshev",
ripple_db=0.1,
first_element="series",
q_l=80,
q_c=120,
)
env = LCBPFOptimizationEnv(
n=n,
LC_base=lc_base,
Q_values=q_values,
f1=f1,
f2=f2,
f0=f0,
fstart=1.5,
fstop=3.5,
points=points,
fq=f0,
z0=50.0,
first_element="series",
targets=default_targets(),
objective="min_insertion_loss",
x_limit=0.5,
action_step=0.05,
max_steps=50,
success_reward=100.0,
reset_noise_scale=reset_noise_scale,
)
return env, q_values, {"n": n, "f0": f0, "f1": f1, "f2": f2}
def make_model(algorithm, env, seed, verbose=1):
algorithm = algorithm.lower()
if algorithm == "ppo":
from stable_baselines3 import PPO
return PPO(
"MlpPolicy",
env,
verbose=verbose,
seed=seed,
learning_rate=3e-4,
n_steps=256,
batch_size=64,
gamma=0.95,
device="cpu",
)
if algorithm == "sac":
from stable_baselines3 import SAC
return SAC(
"MlpPolicy",
env,
verbose=verbose,
seed=seed,
learning_rate=3e-4,
buffer_size=100_000,
learning_starts=500,
batch_size=256,
gamma=0.95,
device="cpu",
)
if algorithm == "td3":
from stable_baselines3 import TD3
from stable_baselines3.common.noise import NormalActionNoise
n_actions = env.action_space.shape[0]
action_noise = NormalActionNoise(
mean=np.zeros(n_actions),
sigma=0.1 * np.ones(n_actions),
)
return TD3(
"MlpPolicy",
env,
action_noise=action_noise,
verbose=verbose,
seed=seed,
learning_rate=1e-3,
buffer_size=100_000,
learning_starts=500,
batch_size=256,
gamma=0.95,
device="cpu",
)
raise ValueError(f"Unsupported algorithm: {algorithm!r}")
def run_training(
algorithm,
total_timesteps=None,
seed=0,
points=801,
reset_noise_scale=0.05,
check=True,
verbose=1,
):
"""学習を実行し、モデル、環境、設計情報を返す。"""
algorithm = algorithm.lower()
if algorithm not in DEFAULT_TIMESTEPS:
raise ValueError(f"Unsupported algorithm: {algorithm!r}")
if total_timesteps is None:
total_timesteps = DEFAULT_TIMESTEPS[algorithm]
if total_timesteps <= 0:
raise ValueError("total_timesteps must be positive.")
env, q_values, design = build_environment(points, reset_noise_scale)
if check:
from stable_baselines3.common.env_checker import check_env
check_env(env, warn=True)
# check_envが試したランダムactionを学習結果に混ぜない。
env.clear_best_result()
model = make_model(algorithm, env, seed=seed, verbose=verbose)
model.learn(total_timesteps=int(total_timesteps))
return model, env, q_values, design
def evaluate_lc_design(lc_elements, q_values, design, name, points=2001):
"""指定したLC値を表示・保存用の高解像度グリッドで再評価する。"""
network = bpf.lossy_bpf(
n=design["n"],
LC_elements=lc_elements,
Q_values=q_values,
fq=design["f0"],
fstart=1.5,
fstop=3.5,
points=points,
z0=50.0,
first_element="series",
name=name,
)
analysis = bpf.evaluate_bpf(
network,
f1=design["f1"],
f2=design["f2"],
f0=design["f0"],
)
return network, analysis
def save_result(path, lc_elements, q_values, x, reward, analysis):
"""LC値、探索変数、報酬、解析値をpickle不要のNPZとして保存する。"""
np.savez(
path,
LC_elements=lc_elements,
Q_values=np.asarray(q_values, dtype=float),
x=np.asarray(x, dtype=float),
reward=float(reward),
best_x=np.asarray(x, dtype=float),
best_reward=float(reward),
analysis_keys=np.asarray(list(analysis)),
analysis_values=np.asarray(list(analysis.values()), dtype=float),
)
def _parser(default_algorithm):
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--algorithm",
choices=sorted(DEFAULT_TIMESTEPS),
default=default_algorithm,
)
parser.add_argument("--timesteps", type=int, default=None)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--points", type=int, default=801)
parser.add_argument("--reset-noise", type=float, default=0.05)
parser.add_argument("--plot", action="store_true", help="応答グラフを画面表示する")
parser.add_argument("--save-prefix", type=Path, default=None)
parser.add_argument("--quiet", action="store_true")
parser.add_argument("--skip-env-check", action="store_true")
return parser
def main(default_algorithm="ppo"):
args = _parser(default_algorithm).parse_args()
# 合成直後のx=0を、reset時のランダム化とは無関係に評価する。
initial_env, initial_q_values, initial_design = build_environment(
points=args.points,
reset_noise_scale=0.0,
)
initial_lc = initial_env.LC_base.copy()
initial_ntwk, initial_analysis = evaluate_lc_design(
initial_lc,
initial_q_values,
initial_design,
name="Initial synthesized BPF",
)
initial_reward, _ = bpf_score.score_bpf_analysis(
initial_analysis,
targets=initial_env.targets,
objective=initial_env.objective,
)
print("\n=== Initial design (before training) ===")
print(f"Initial reward: {initial_reward:.6g}")
bpf.print_lc_table(initial_lc, initial_q_values)
bpf.print_analysis(initial_analysis)
initial_plot_path = None
if args.save_prefix is not None:
args.save_prefix.parent.mkdir(parents=True, exist_ok=True)
save_result(
str(args.save_prefix) + "_initial_result.npz",
initial_lc,
initial_q_values,
np.zeros(initial_env.dim),
initial_reward,
initial_analysis,
)
initial_plot_path = str(args.save_prefix) + "_initial_response.png"
if args.plot or initial_plot_path is not None:
if args.plot:
print("Close the initial response window to start training.")
bpf.plot_bpf_response(
initial_ntwk,
f1=initial_design["f1"],
f2=initial_design["f2"],
f0=initial_design["f0"],
analysis=initial_analysis,
title=initial_ntwk.name,
save_path=initial_plot_path,
show=args.plot,
)
model, env, q_values, design = run_training(
algorithm=args.algorithm,
total_timesteps=args.timesteps,
seed=args.seed,
points=args.points,
reset_noise_scale=args.reset_noise,
check=not args.skip_env_check,
verbose=0 if args.quiet else 1,
)
best = env.get_best_result()
best_lc = best["best_LC_elements"]
if best_lc is None:
raise RuntimeError("Training completed without evaluating a filter.")
best_ntwk, best_analysis = evaluate_lc_design(
best_lc,
q_values,
design,
name=f"{args.algorithm.upper()} optimized BPF",
)
print("\n=== Optimized design (after training) ===")
print(f"Best training reward: {best['best_reward']:.6g}")
bpf.print_lc_table(best_lc, q_values)
bpf.print_analysis(best_analysis)
plot_path = None
if args.save_prefix is not None:
args.save_prefix.parent.mkdir(parents=True, exist_ok=True)
model.save(str(args.save_prefix) + "_model")
save_result(
str(args.save_prefix) + "_result.npz",
best_lc,
q_values,
best["best_x"],
best["best_reward"],
best_analysis,
)
plot_path = str(args.save_prefix) + "_response.png"
if args.plot or plot_path is not None:
bpf.plot_bpf_response(
best_ntwk,
f1=design["f1"],
f2=design["f2"],
f0=design["f0"],
analysis=best_analysis,
title=best_ntwk.name,
save_path=plot_path,
show=args.plot,
)
return best
|
PPOを使う場合は
もともと:
学習結果:
ちょっとフラットになった感じか。別のアルゴリズム試してみよう。
« RF Weekly Digest 2026/8/17-8/26 (Codex(GPT-5.6 Sol)によるRF情報週刊まとめ) | トップページ | 高周波・RFニュース 2026年8月25日 Microwave Journal8月スペシャルフォーカスはミリ波、everythingRFがLTE Cat 1bisのeBook発行、NordicのnRF93M1はCat 1bis対応、BroadcomがTri-band WI-Fi 8 SoCについて解説など »
「パソコン・インターネット」カテゴリの記事
「学問・資格」カテゴリの記事
「日記・コラム・つぶやき」カテゴリの記事
« RF Weekly Digest 2026/8/17-8/26 (Codex(GPT-5.6 Sol)によるRF情報週刊まとめ) | トップページ | 高周波・RFニュース 2026年8月25日 Microwave Journal8月スペシャルフォーカスはミリ波、everythingRFがLTE Cat 1bisのeBook発行、NordicのnRF93M1はCat 1bis対応、BroadcomがTri-band WI-Fi 8 SoCについて解説など »




コメント