如何将函数作为参数传入另一函数?附RRT算法代码问询
Got it, let's walk through how to do this—Python makes passing functions as arguments super straightforward since functions are first-class citizens here. Here's a step-by-step breakdown tailored to your RRT code:
1. 先定义你的距离计算函数
First, let's assume your Calculate_Distance looks something like this (adjust based on your actual coordinate system, e.g., 2D/3D space):
def Calculate_Distance(point_a, point_b): # 示例:计算欧氏距离(2D) return ((point_a[0] - point_b[0])**2 + (point_a[1] - point_b[1])**2)**0.5
2. 修改Initiate_Sampling以接收距离函数参数
Update your Initiate_Sampling function to accept a distance_func parameter. Inside the function, you'll use this passed-in function to compute the distance between points, then compare it against EPS:
import random # 假设你的采样需要生成随机点 def Initiate_Sampling(start_point, eps, distance_func): # 模拟生成随机采样点(根据你的RRT逻辑调整这部分) random_sample = (random.uniform(0, 10), random.uniform(0, 10)) # 使用传入的距离函数计算起点和采样点的距离 sampled_distance = distance_func(start_point, random_sample) # 和EPS比较,执行你的RRT采样逻辑 if sampled_distance < eps: print(f"采样点符合要求,距离起点:{sampled_distance:.2f}") return random_sample else: print(f"采样点距离过大,跳过该点") return None
3. 调用时直接传递函数名
When calling Initiate_Sampling, just pass the Calculate_Distance function name (don't add parentheses—that would execute the function immediately and pass its return value instead):
# 示例调用 start = (5, 5) EPS = 2.0 # 你的阈值 result = Initiate_Sampling(start, EPS, Calculate_Distance) print(f"最终采样得到的点:{result}")
关键注意事项
- 不要加括号: When passing the function, use
Calculate_Distanceinstead ofCalculate_Distance(). The latter runs the function and passes its output, not the function itself. - 解耦优势: This approach decouples your sampling logic from the distance calculation. If you later want to switch to a different distance metric (like Manhattan distance), you just need to define a new distance function and pass it in—no changes needed to
Initiate_Sampling.
扩展:如果距离函数需要额外参数
If your Calculate_Distance requires extra parameters (e.g., for 3D space or weighted distances), you can wrap it with a lambda to fix those parameters:
# 带额外参数的距离函数 def Calculate_Distance(point_a, point_b, dimensions): return sum((p1 - p2)**2 for p1, p2 in zip(point_a, point_b))**0.5 # 调用时用lambda固定dimensions参数 result = Initiate_Sampling(start, EPS, lambda p1, p2: Calculate_Distance(p1, p2, dimensions=3))
内容的提问来源于stack exchange,提问作者Adwait

