基于afsim的训练多个智能体算法控制(持续关注和收藏后续会同步整个源码包)

python 复制代码
import os
import math
import numpy as np

# def get_real_obs(obs_json):
#     # 解析obs
#     obs_array = []
#     for entity in obs_json['entityStatusData']:
#         index = int(entity['Index'])
#         lon = float(entity['Lon'])
#         lat = float(entity['Lat'])
#         # alt = float(entity['Alt'])
#         obs_array.extend([index,lon,lat])
#     obs_array = np.array(obs_array,dtype=np.float32)
#     return obs_array
# point 24:30:03.51n 119:02:27.56e // Top Right
#       point 24:46:32.35n 119:19:50.98e // Bottom Right
#       point 24:36:48.89n 119:30:55.89e // Bottom Left
#       point 24:20:12.71n 119:13:51.33e // Top Left


class Point:
    def __init__(self, lon, lat):
        self.lon = lon
        self.lat = lat


def dms_to_dds(dms_lat, dms_lon):
    """Converts coordinates from DMS to DD format."""
    try:
        lat_dms = dms_lat.split(':')
        lon_dms = dms_lon.split(':')

        if len(lat_dms) != 3 or len(lon_dms) != 3:
            return None
        lat_dd = float(lat_dms[0]) + float(lat_dms[1]) / 60 + float(lat_dms[2].replace('n', '').replace('s', '')) / 3600
        lon_dd = float(lon_dms[0]) + float(lon_dms[1]) / 60 + float(lon_dms[2].replace('e', '').replace('w', '')) / 3600
    
        if 's' in dms_lat.lower():
            lat_dd *= -1
        if 'w' in dms_lon.lower():
            lon_dd *= -1
        # print(lon_dd, lat_dd)
        return Point(lon_dd, lat_dd)
        # return [lon_dd, lat_dd]
    except (ValueError, IndexError):
        return None

def dms_to_dd(dms_lat, dms_lon):
  """
    将坐标从DMS转换为DD格式。
    Args:
    dms_lat:dms格式的纬度(例如,"24:32:13.08")。
    dms_lon:dms格式的经度(例如,"119:08:45.57")。
    return:
    包含DD格式的纬度和经度的元组(例如,(24.53691119.145992))。
    如果输
  """
  try:
    lat_dms = dms_lat.split(':')
    lon_dms = dms_lon.split(':')

    if len(lat_dms) != 3 or len(lon_dms) != 3:
      return None

    lat_dd = float(lat_dms[0]) + float(lat_dms[1]) / 60 + float(lat_dms[2]) / 3600
    lon_dd = float(lon_dms[0]) + float(lon_dms[1]) / 60 + float(lon_dms[2]) / 3600

    return [lat_dd, lon_dd]
    # return Point(lon_dd, lat_dd)
  except (ValueError, IndexError):
    return None  # Handle cases with invalid input format


def get_task_point():
    A_lat = '24:30:03.51' 
    A_lon = '119:02:27.56'
    B_lat = '24:46:32.3'
    B_lon = '119:19:50.98'
    C_lat = '24:36:48.89' 
    C_lon = '119:30:55.89'
    D_lat = '24:20:12.71'
    D_lon = '119:13:51.33'
    
    A_point = dms_to_dd(A_lat,A_lon)
    B_point = dms_to_dd(B_lat,B_lon)
    C_point = dms_to_dd(C_lat,C_lon)
    D_point = dms_to_dd(D_lat,D_lon)
    #处理DMS格式无效的情况
    if any(p is None for p in [A_point, B_point, C_point, D_point]):
        return None
    # print(A_point,B_point,C_point,D_point)
    return A_point,B_point,C_point,D_point

# 计算奖励用到的点
def get_base_point():
    A_lat = '24:30:03.51' 
    A_lon = '119:02:27.56'
    B_lat = '24:46:32.3'
    B_lon = '119:19:50.98'
    C_lat = '24:36:48.89' 
    C_lon = '119:30:55.89'
    D_lat = '24:20:12.71'
    D_lon = '119:13:51.33'
    A_point = dms_to_dd(A_lat,A_lon)
    B_point = dms_to_dd(B_lat,B_lon)
    C_point = dms_to_dd(C_lat,C_lon)
    D_point = dms_to_dd(D_lat,D_lon)
    center_lat_point = (A_point[0]+B_point[0]+C_point[0]+D_point[0])/4
    center_lon_point = (A_point[1]+B_point[1]+C_point[1]+D_point[1])/4
    return A_point,B_point,C_point,D_point,center_lat_point,center_lon_point

def get_danger_point():
    """
    24:26:18.92n 119:15:25.68e
    24:32:09.91n 119:08:51.77e
    24:25:08.29n 119:07:45.31e
    24:20:05.78n 119:13:50.01e
    """
    A_lat = '24:26:18.92' 
    A_lon = '119:15:25.68'
    B_lat = '24:32:09.91'
    B_lon = '119:08:51.77'
    C_lat = '24:25:08.29' 
    C_lon = '119:07:45.31'
    D_lat = '24:20:05.78'
    D_lon = '119:13:50.01'
    A_point = dms_to_dd(A_lat,A_lon)
    B_point = dms_to_dd(B_lat,B_lon)
    C_point = dms_to_dd(C_lat,C_lon)
    D_point = dms_to_dd(D_lat,D_lon)
    center_lat_point = (A_point[0]+B_point[0]+C_point[0]+D_point[0])/4
    center_lon_point = (A_point[1]+B_point[1]+C_point[1]+D_point[1])/4
    return A_point,B_point,C_point,D_point,center_lat_point,center_lon_point



def get_real_obs(obs_json):

    # 解析obs
    # 初始化obs_array
    done = False
    obs_array = []
    if obs_json == {}:
        done = True
        return done,obs_array
    # 获取实体数据
    if len(obs_json['entityStatusData']) < 2:
        done = True
        return done,obs_array
    entity_A = obs_json['entityStatusData'][0]  # 第一个实体
    entity_B = obs_json['entityStatusData'][1]  # 第二个实体
    
    # 提取实体A的相关信息
    index_A = int(entity_A['Index'])
    lat_A = float(entity_A['Lat'])
    lon_A = float(entity_A['Lon'])
    
    # 提取实体B的相关信息
    index_B = int(entity_B['Index'])
    lat_B = float(entity_B['Lat'])
    lon_B = float(entity_B['Lon'])
    
    # 计算经纬度差
    lat_diff = lat_B - lat_A
    lon_diff = lon_B - lon_A
    task_points = get_task_point()
    if task_points is None:
        return True, [] #如果DMS转换出错,返回空数组
    # 展平列表
    flattened_task_points = [item for sublist in task_points for item in sublist]

    for i in range(len(flattened_task_points)):
        if i % 2 == 0:
            flattened_task_points[i] = (flattened_task_points[i] - lat_A)/2
        else:
            flattened_task_points[i] = (flattened_task_points[i] - lon_A)/2
    
    # 将lat_A lon_A 和lat_B lon_B标准化
    lat_A = (lat_A - 23.5)/2
    lon_A = (lon_A - 118.5)/2
    lat_B = (lat_B - 23.5)/2
    lon_B = (lon_B - 118.5)/2
    # 0112 加入红方实体信息
    # 将所有数据合并到obs_array中。注意这里需要类型检查
    obs_array.extend([float(lat_A), float(lon_A),float(lat_B),float(lon_B), *flattened_task_points])

    # 转换为NumPy数组并返回
    obs_array = np.array(obs_array, dtype=np.float32)
    
    return done,obs_array

# 打印途径点
def print_point(action):
        for ac in action['command']['wayPoints']:
            lat = float(ac['Lat'])
            lon = float(ac['Lon'])
            d,m,s = decimal_to_dms(lat)
            dd,mm,ss = decimal_to_dms(lon)
            print(f"lat:{d}:{m}:{s},lon:{dd}:{mm}:{ss}")


# 获取敌方态势信息
def get_enemy_data(obs_all,path):
    file_name = 'obs_enemy_data.txt'
    file_path = os.path.join(path, file_name)
    with open(file_path, 'w') as file:
        print("")
    for obs_all in obs_all:
        # if obs_all['ID'] in ['3', '4']:
            # 构建所需的格式化字符串
            data_str = f"ID: {obs_all['ID']}, pitchNED: {obs_all['pitchNED']}, rollNED: {obs_all['rollNED']}, headingNED: {obs_all['headingNED']}\n"
            # 将数据追加写入指定文件
            with open(file_path, 'a') as file:
                file.write(data_str)
    return True

# 小数度转经纬度
def decimal_to_dms(decimal_degree):
    # 将小数度(Decimal Degrees)转换为度分秒(DMS)
    # 获取度(整数部分)
    degree = int(decimal_degree)
    # 获取分钟(去掉度之后的部分,乘以 60)
    minute = int((decimal_degree - degree) * 60)
    # 获取秒(去掉分钟之后的部分,乘以 60)
    second = round((decimal_degree - degree - minute / 60) * 3600, 4)
    return degree,minute,second

# 经纬度转小数度
def dms_to_decimal(degree,minute,second):
    # 将度分秒(DMS)转换为小数度(Decimal Degrees, DD)
    decimal = degree + minute / 60 + second / 3600
    return decimal

def normalize_state(state, lon_min=118.5, lon_max=120.5, lat_min=23.5, lat_max=25.5,k=1):
    """
    Normalizes the state array containing latitude and longitude values.
    规范化包含纬度和经度值的状态数组
    Args:
        state: A NumPy array of shape (N,) where N is even, containing alternating latitude and longitude values.
        lon_min: Minimum longitude value for normalization.
        lon_max: Maximum longitude value for normalization.
        lat_min: Minimum latitude value for normalization.
        lat_max: Maximum latitude value for normalization.

    Returns:
        A NumPy array with normalized latitude and longitude values, or None if input is invalid.  
        Values outside the specified range will be clipped to the range's boundaries.
    """

    if len(state) % 2 != 0:
        print("Error: State array must have an even number of elements (latitude-longitude pairs).")
        return None

    latitudes = state[::2]  # Extract latitudes
    longitudes = state[1::2] # Extract longitudes

    #Clip values to the specified range
    latitudes = np.clip(latitudes, lat_min, lat_max)
    longitudes = np.clip(longitudes, lon_min, lon_max)


    # Normalize latitudes and longitudes separately
    normalized_latitudes = (latitudes - lat_min) / (lat_max - lat_min)
    normalized_longitudes = (longitudes - lon_min) / (lon_max - lon_min)

    # 将标准化的纬度和经度交错
    normalized_state = np.empty_like(state, dtype=np.float32)
    normalized_state[::2] = normalized_latitudes
    normalized_state[1::2] = normalized_longitudes

    # 对实体位置着重处理
    normalized_state = np.hstack([normalized_state[:2] * k, normalized_state[2:]])

    return normalized_state

# 将状态归一化处理
def normalize_state_v1(state):
    # 假设已知的经纬度和高度的范围
    longitude_range = (-180, 180)  # 经度范围
    latitude_range = (-90, 90)    # 纬度范围
    altitude_range = (0, 15000)   # 高度范围
    position_range = (-4, 4)

    # 获取 ID 和经纬度高度数据
    agent_id_1 = state[0]  # 第0个是ID
    longitude_1 = state[1] 
    latitude_1 = state[2]
    # altitude_1 = state[3]
    
    agent_id_2 = state[3]  # 第4个是ID
    longitude_2 = state[4]
    latitude_2 = state[5]
    # altitude_2 = state[7]

    posi_1 = state[6]
    posi_2 = state[7]
    
    # 对经度进行归一化
    normalized_longitude_1 = (longitude_1 - longitude_range[0]) / (longitude_range[1] - longitude_range[0])
    normalized_longitude_2 = (longitude_2 - longitude_range[0]) / (longitude_range[1] - longitude_range[0])
    
    # 对纬度进行归一化
    normalized_latitude_1 = (latitude_1 - latitude_range[0]) / (latitude_range[1] - latitude_range[0])
    normalized_latitude_2 = (latitude_2 - latitude_range[0]) / (latitude_range[1] - latitude_range[0])
    
    # 对高度进行归一化
    # normalized_altitude_1 = (altitude_1 - altitude_range[0]) / (altitude_range[1] - altitude_range[0])
    # normalized_altitude_2 = (altitude_2 - altitude_range[0]) / (altitude_range[1] - altitude_range[0])

    # 对相对位置进行归一化
    norm_lat_diff = (posi_1 - position_range[0]) / (position_range[1] - position_range[0])
    norm_lon_diff = (posi_2 - position_range[0]) / (position_range[1] - position_range[0])
    
    # 返回归一化后的状态
    normalized_state = np.array([agent_id_1, normalized_longitude_1, normalized_latitude_1,
                                 agent_id_2, normalized_longitude_2, normalized_latitude_2,
                                 norm_lat_diff, norm_lon_diff])
    
    return normalized_state

def generate_waypoints_twopoint(lat, lon, alt,action, num_segments=15):
    # 起点到中点的分段(action[0] 和 action[1] 控制第一个部分)
    lat1 = float(lat)
    lon1 = float(lon)
    lat_mid = lat1 + action[0]  # 终点的经度
    lon_mid = lon1 + action[1]  # 终点的纬度
    
    # 中点到终点的分段(action[2] 和 action[3] 控制第二个部分)
    lat_end = lat_mid + action[2]
    lon_end = lon_mid + action[3]
    
    # 使用np.linspace进行插值:从起点到中点,以及从中点到终点,分成num_segments个小段
    waypoints = []

    # 插值生成从起点到中点的路径
    for i in range(num_segments):
        t = i / (num_segments - 1)  # 线性插值系数
        lat_new = lat1 + t * (lat_mid - lat1)
        lon_new = lon1 + t * (lon_mid - lon1)
        waypoints.append({"Lat": str(lat_new), "Lon": str(lon_new), "Alt": str(alt), "Speed": "150"})
    
    # 插值生成从中点到终点的路径
    for i in range(num_segments):
        t = i / (num_segments - 1)  # 线性插值系数
        lat_new = lat_mid + t * (lat_end - lat_mid)
        lon_new = lon_mid + t * (lon_end - lon_mid)
        waypoints.append({"Lat": str(lat_new), "Lon": str(lon_new), "Alt": str(alt), "Speed": "600"})

    return waypoints

def generate_waypoints_onepoint(idx,speed,lat, lon, alt,action, num_segments=15):
    # 起点到终点的分段(action[0] 和 action[1] 控制第一个部分)
    lat1 = float(lat)
    lon1 = float(lon)
    lon_mid = lon1 + action[0]  # 终点的纬度
    lat_mid = lat1 + action[1]# 终点的经度
    # 使用np.linspace进行插值:从起点到中点,以及从中点到终点,分成num_segments个小段
    waypoints_last_step = []
    waypoints_last_step.append({"Lat": str(lat_mid), "Lon": str(lon_mid), "Alt": str(alt), "Speed": str(0)})
    waypoints = []
    waypoints.append({"Lat": str(lat_mid), "Lon": str(lon_mid), "Alt": str(alt), "Speed": str(speed[idx])})
    # 插值生成从起点到中点的路径
    # for i in range(num_segments):
    #     t = i / (num_segments - 1)  # 线性插值系数
    #     lat_new = lat1 + t * (lat_mid - lat1)
    #     lon_new = lon1 + t * (lon_mid - lon1)
    #     waypoints.append({"Lat": str(lat_new), "Lon": str(lon_new), "Alt": str(alt), "Speed": "150"})
    
    
    return waypoints,waypoints_last_step

# 经纬度替换
def replace_lat_lon(state,lat,lon):
    state[1] = lon
    state[2] = lat
    return state


def get_cmd_actions(obs,index,action):
    action_speed = [150,150,150,150,150]
    lat = obs['entityStatusData'][0]['Lat']
    Lon = obs['entityStatusData'][0]['Lon']
    Alt = obs['entityStatusData'][0]['Alt']
    waypoints,wayp_ls_step = generate_waypoints_onepoint(index,action_speed,lat, Lon,Alt, action, num_segments=1)
    move_cmd = {
            "command":{
                "ID":"1",
                "Type": "route",
                "wayPoints":waypoints
            }
        }
    move_cmd_ls_step = {
            "command":{
                "ID":"1",
                "Type": "route",
                "wayPoints":wayp_ls_step
            }
        }
    return move_cmd,move_cmd_ls_step



def euclidean_distance(lat_A, lon_A, lat_B, lon_B):
    return math.sqrt((lat_B - lat_A) ** 2 + (lon_B - lon_A) ** 2)

#根据二维平面坐标计算奖励(忽略曲率)
def calculate_reward(point_A,point_B, epsilon=0.05, k=1, stability_bonus=0.1):
    lat_A = float(point_A['Lat'])
    lon_A = float(point_A['Lon'])
    lat_B = float(point_B['Lat'])
    lon_B = float(point_B['Lon'])
    # 计算两点之间的欧几里得距离
    distance = euclidean_distance(lat_A, lon_A, lat_B, lon_B)
    
    # 如果距离小于等于阈值 epsilon,奖励为 1,并且加上稳定性奖励
    if distance <= epsilon:
        reward = 1 + stability_bonus
    else:
        # 否则,奖励为 1 减去距离的惩罚项
        reward =  - k * distance
    
    return reward

# 根据度计算距离
def _cal_dis(point_A,point_B):
    R = 6371.0
    lat_A = float(point_A['Lat'])
    lon_A = float(point_A['Lon'])
    lat_B = float(point_B['Lat'])
    lon_B = float(point_B['Lon'])
    # 经纬度转为弧度
    lat1 = math.radians(lat_A)
    lon1 = math.radians(lon_A)
    lat2 = math.radians(lat_B)
    lon2 = math.radians(lon_B)
    # 计算经纬度差值
    dlat = lat2 - lat1
    dlon = lon2 - lon1
    # Haversine公式计算球面距离
    a = math.sin(dlat / 2)**2 + math.cos(lat1) * math.cos(lat2) * math.sin(dlon / 2)**2
    c = 2 * math.atan2(math.sqrt(a), math.sqrt(1 - a))
    
    # 计算距离
    distance = R * c  # 距离单位:公里
    return -distance
def cal_dis(point_A,point_B):

    lat_A = float(point_A['Lat'])
    lon_A = float(point_A['Lon'])
    lat_B = float(point_B['Lat'])
    lon_B = float(point_B['Lon'])
    distance = math.sqrt((lat_A - lat_B) ** 2 + (lon_A - lon_B) ** 2)
    return -distance

def cal_location_reward_out(blue_point_obs,point_base):
    A_p,B_p,C_p,D_p,Cent_lat_p,Cent_lon_p = point_base
    lat_A = float(blue_point_obs['Lat'])
    lon_A = float(blue_point_obs['Lon'])
    print(f"lat_A:{lat_A},lon_A:{lon_A}")
    # 计算欧几里得距离
    distance_A = euclidean_distance(lat_A, lon_A, A_p[0], A_p[1])
    distance_B = euclidean_distance(lat_A, lon_A, B_p[0], B_p[1])
    distance_C = euclidean_distance(lat_A, lon_A, C_p[0], C_p[1])
    distance_D = euclidean_distance(lat_A, lon_A, D_p[0], D_p[1])
    reward_location = -50*(distance_A+distance_B+distance_C+distance_D)
    # reward_center = -50*euclidean_distance(lat_A, lon_A, Cent_lat_p, Cent_lon_p)# 在活动区域外,给一个较大的负奖励
    return reward_location

def cal_location_reward_in(blue_point_obs,point_base):
    A_p,B_p,C_p,D_p,Cent_lat_p,Cent_lon_p = point_base
    lat_A = float(blue_point_obs['Lat'])
    lon_A = float(blue_point_obs['Lon'])
    reward_center = 10*euclidean_distance(lat_A, lon_A, Cent_lat_p, Cent_lon_p)# 在活动区域内给正奖励# 在危险区给负奖励
    return reward_center
def cal_location_reward_out_danger(blue_point_obs,point_base):
    A_p,B_p,C_p,D_p,Cent_lat_p,Cent_lon_p = point_base
    lat_A = float(blue_point_obs['Lat'])
    lon_A = float(blue_point_obs['Lon'])
    reward_center = 2*euclidean_distance(lat_A, lon_A, Cent_lat_p, Cent_lon_p)# 在活动区域内给正奖励# 在危险区给负奖励
    return reward_center
def cal_location_reward_in_danger(blue_point_obs,point_base):
    A_p,B_p,C_p,D_p,Cent_lat_p,Cent_lon_p = point_base
    lat_A = float(blue_point_obs['Lat'])
    lon_A = float(blue_point_obs['Lon'])
    reward_center = -20*euclidean_distance(lat_A, lon_A, Cent_lat_p, Cent_lon_p)# 在活动区域内给正奖励# 在危险区给负奖励
    return reward_center
# 任务区
def is_point_in_polygon(Point):
    """射线算法检查一个点是否在多边形内"""
    # 矩形坐标 (顺时针)
    points = [
    "24:30:03.51n 119:02:27.56e",  # Top Right
        "24:46:32.35n 119:19:50.98e",  # Bottom Right
        "24:36:48.89n 119:30:55.89e",  # Bottom Left
        "24:20:12.71n 119:13:51.33e"   # Top Left
    ]

    # Convert DMS to DD
    polygon_points = []
    for point in points:
        lat_str, lon_str = point.split(' ')
        dd_point = dms_to_dds(lat_str,lon_str)
        if dd_point:
            polygon_points.append(dd_point)
        else:
            print("Error: Invalid DMS format.")
    x = float(Point['Lon'])
    y = float(Point['Lat'])
    n = len(polygon_points)
    inside = False
    p1 = polygon_points[0]
    for i in range(1, n + 1):
        p2 = polygon_points[i % n]
        if y > min(p1.lat, p2.lat):
            if y <= max(p1.lat, p2.lat):
                if x <= max(p1.lon, p2.lon):
                    if p1.lat != p2.lat:
                        xinters = (y - p1.lat) * (p2.lon - p1.lon) / (p2.lat - p1.lat) + p1.lon
                        if xinters >= x:
                            inside = not inside
        p1 = p2

    if inside:
        print(f"点{Point}在区域内")
    else:
        print(f"点{Point}在区域外")
    return inside
# 危险区
def is_point_in_polygon_danger(Point):
    """射线算法检查一个点是否在多边形内"""
    # 矩形坐标 (顺时针)
    points = [
    "24:26:18.92n 119:15:25.68e",  # Top Right
        "24:32:09.91n 119:08:51.77e",  # Bottom Right
        "24:25:08.29n 119:07:45.31e",  # Bottom Left
        "24:20:05.78n 119:13:50.01e"   # Top Left
    ]

    # Convert DMS to DD
    polygon_points = []
    for point in points:
        lat_str, lon_str = point.split(' ')
        dd_point = dms_to_dds(lat_str,lon_str)
        if dd_point:
            polygon_points.append(dd_point)
        else:
            print("Error: Invalid DMS format.")
    x = float(Point['Lon'])
    y = float(Point['Lat'])
    n = len(polygon_points)
    inside = False
    p1 = polygon_points[0]
    for i in range(1, n + 1):
        p2 = polygon_points[i % n]
        if y > min(p1.lat, p2.lat):
            if y <= max(p1.lat, p2.lat):
                if x <= max(p1.lon, p2.lon):
                    if p1.lat != p2.lat:
                        xinters = (y - p1.lat) * (p2.lon - p1.lon) / (p2.lat - p1.lat) + p1.lon
                        if xinters >= x:
                            inside = not inside
        p1 = p2

    if inside:
        print(f"点{Point}在危险区域内")
    else:
        print(f"点{Point}在危险区域外")
    return inside

# points = {"Lon":119.5,"Lat":24.5}
# ind = is_point_in_polygon(points)
# print(ind)
# points = {"Lon":119.5,"Lat":24.6}
# ind = is_point_in_polygon(points)
# print(ind)

# 计算奖励
"""
    A_lat = '24:30:03.51' 
    A_lon = '119:02:27.56'
    B_lat = '24:46:32.3'
    B_lon = '119:19:50.98'
    C_lat = '24:36:48.89' 
    C_lon = '119:30:55.89'
    D_lat = '24:20:12.71'
    D_lon = '119:13:51.33'
    A_point: [24.500975, 119.040985]
    B_point: [24.775639, 119.330826]
    C_point: [24.61358, 119.515526]
    D_point: [24.336864, 119.23093]
    """
def cal_reward(obs,action,max_lat=24.775639,min_lat=24.336864,max_lon=119.515526,min_lon=119.040985):
    is_inorout = False
    reward_win = 0
    # 判断仿真中还有几个实体
    if len(obs['entityStatusData'])<2:
        done = True
    else:
        done = False
    # 红方位置
    red_point_obs = obs['entityStatusData'][-1]
    # 蓝方位置
    blue_point_obs = obs['entityStatusData'][0]
    # 蓝方要移动的目标位置
    newblue_point_obs = action['command']['wayPoints'][-1]

    # 蓝方纬度和经度
    blue_lat = float(blue_point_obs['Lat'])
    blue_lon = float(blue_point_obs['Lon'])
    # 任务区区域
    point_base = get_base_point()
    # 危险区区域
    point_danger = get_danger_point()

    # 任务区奖励
    is_inside_Missionarea = is_point_in_polygon(newblue_point_obs)
    if is_inside_Missionarea:
        reward_location_mission = cal_location_reward_in(newblue_point_obs,point_base)
    else:
        reward_location_mission = cal_location_reward_out(newblue_point_obs,point_base)

    # 危险区惩罚
    is_inside_dangerarea = is_point_in_polygon_danger(newblue_point_obs)
    if is_inside_dangerarea:
        reward_location_danger = cal_location_reward_in_danger(newblue_point_obs,point_danger)
    else:
        reward_location_danger = cal_location_reward_out_danger(newblue_point_obs,point_danger)
    # reward = cal_dis(red_point_obs,newblue_point_obs)
    # reward = calculate_reward(red_point_obs,newblue_point_obs)
    current_dis = cal_dis(blue_point_obs,newblue_point_obs)

    # 胜利奖励
    if done and obs['entityStatusData']['index']=='1':
        reward_win = 500
    elif done:
        reward_win = -500

    # 边界惩罚(判断实体位置是否在指定位置内(超出固定区域,该Episode结束))
    lat_blue = float(obs['entityStatusData'][0]['Lat'])
    lon_blue = float(obs['entityStatusData'][0]['Lon'])
    if lat_blue >25.5 or lat_blue < 23.5 or  lon_blue > 120.5 or lon_blue < 118.5:
        done = True
        reward_edge = -100
    else:
        reward_edge = 0
    reward = reward_location_mission + reward_location_danger + reward_win + reward_edge
    print(f"目标位置:{newblue_point_obs}")
    return reward,done,newblue_point_obs,current_dis

# 计算延时
def cal_time_sleep(dis):
    times = abs(dis)*1000/150
    print(f"延时大小:{times}")
    return times




def cal_dis_lonandlat(point_A,point_B):
    R = 6371.0
    lat_A = point_A
    lon_A = point_A
    lat_B = point_B
    lon_B = point_B
    # 经纬度转为弧度
    lat1 = math.radians(lat_A)
    lon1 = math.radians(lon_A)
    lat2 = math.radians(lat_B)
    lon2 = math.radians(lon_B)
    # 计算经纬度差值
    dlat = lat2 - lat1
    dlon = lon2 - lon1
    # Haversine公式计算球面距离
    a = math.sin(dlat / 2)**2 + math.cos(lat1) * math.cos(lat2) * math.sin(dlon / 2)**2
    c = 2 * math.atan2(math.sqrt(a), math.sqrt(1 - a))
    
    # 计算距离
    distance = R * c  # 距离单位:公里
    return -distance
# 
# print(cal_dis_lonandlat(0.0,0.00135))
def take_obs(obs):
    flag = False
    entity_A = obs['entityStatusData'][0]  # 第一个实体
    # 提取实体A的经纬
    lat_A = float(entity_A['Lat'])
    lon_A = float(entity_A['Lon'])
    if lat_A>25.2 or lat_A<24 or lon_A >121 or lon_A<118:
        flag = True
    return flag

def process_action(ac, ind=0):
    # 定义允许的 ac 值的规则
    valid_actions = {
        0: [1, 3],  # ind = 0 时,ac 可取 1 或 3
        1: [0, 2],  # ind = 1 时,ac 可取 0 或 2
        2: [1, 3],  # ind = 2 时,ac 可取 1 或 3
        3: [0, 2]   # ind = 3 时,ac 可取 0 或 2
    }
    
    # 如果 ac 不在 valid_actions[ind] 中,则强制将 ac 设为 ind
    if ac not in valid_actions[ind]:
        return ind
    else:
        return ac
    
def conversion_action(ac_index, current_angle, step_size):
    """
    根据给定的动作索引和当前角度计算智能体的下一个动作和方向。

    参数:
    - ac_index (int): 强化学习输出的离散动作索引,0表示保持当前方向,1表示向左偏移45度,2表示向右偏移45度。
    - current_angle (float): 当前的移动方向,以角度表示(0到360度之间)。
    - step_size (float): 每个时间步的移动距离(经度/纬度变化量)。

    返回:
    - new_angle (float): 更新后的移动角度。
    - action (list): 当前选择的动作映射值,包含经度和纬度的变化量。
    """
    
    # 根据动作索引选择新的方向
    possible_angles = [current_angle-15, current_angle,current_angle + 15]
    
    possible_angles = [(angle + 360) % 360 for angle in possible_angles]  # 保证角度在0到360之间

    # 偏移角度列表
    angle_offsets = [-30, -15, 0, 15, 30]
    
    # 获取新的角度
    new_angle = (current_angle + possible_angles[ac_index]) % 360  # 保证角度在 0 到 360 之间


    # 根据 ac_index 选择动作对应的角度
    new_angle = possible_angles[ac_index]

    # 计算新的经纬度增量
    angle_rad = math.radians(new_angle)  # 转换为弧度
    delta_lon = math.cos(angle_rad) * step_size
    delta_lat = math.sin(angle_rad) * step_size
    # 计算新的平面坐标增量(x, y)
    delta_x = step_size * (new_angle / 90)  # 将角度转换为对应的 x 方向变化量
    delta_y = step_size * (new_angle / 90)  # 将角度转换为对应的 y 方向变化量

    # 返回新的角度和动作映射值
    action = [delta_lon, delta_lat]
    # action = [delta_x, delta_y]
    return new_angle, action
相关推荐
保加利亚的风1 小时前
Docker 学习文档(Mac + Docker Desktop 版)
前端·后端
刘名喜1 小时前
第30篇-Spring-Security-7核心概念
后端·kotlin·springboot
明月_清风1 小时前
🚀 AI Agent 完全入门指南:从 LLM 到生产落地,新手必懂的 34 个核心概念
前端·后端·ai编程
程序员阿明2 小时前
spring boot3访问resources下的文件,使得在浏览器url中可以直接访问
java·spring boot·后端
掘金者阿豪2 小时前
MongoDB迁移不想重写代码?一次国产数据库替换踩坑记录
后端
Vec‑Jie2 小时前
MerchantOps-KBQA 实践(十四):本地运行、性能限制与可验证的交付边界
人工智能·后端·python·flask
qo_tn2 小时前
微信机器人-webhook技术文档_02-部署前准备与环境规划
后端
雪隐2 小时前
WPF + MVVM 实战系列02-告别 INPC 手写时代,做个体面的现代 WPF 人
前端·后端·c#