import numpy as np
import matplotlib.pyplot as plt
from PIL import Image
from scipy.signal import convolve, find_peaks

# --- 中文字体 ---
plt.rcParams['font.sans-serif'] = ['Microsoft YaHei'] 
plt.rcParams['axes.unicode_minus'] = False
# --------------------

# ==============================================================================
# === 可调节参数 ===
# 谱线识别灵敏度阈值 (0.0 到 1.0之间, 建议 0.05-0.3)
# 数值越高，表示只识别越深的谱线，灵敏度越低。
LINE_SENSITIVITY_THRESHOLD = 0.08

# 谱线匹配容差 (单位: 像素)
# 如果两条谱线的像素位置差在此范围内，则认为它们是匹配的。
MATCHING_TOLERANCE_PIXELS = 10
# ==============================================================================

def normalize(data):
    """对一维数组进行最小-最大归一化处理"""
    min_val = np.min(data)
    max_val = np.max(data)
    if max_val == min_val:
        return np.zeros_like(data)
    return (data - min_val) / (max_val - min_val)

def get_spectrum_from_image(image_path, target_length=None):
    """从图像加载并提取一维光谱序列"""
    img = Image.open(image_path).convert('L')
    data = np.array(img)
    spectrum = np.mean(data, axis=0)
    
    if target_length and len(spectrum) != target_length:
        if len(spectrum) > target_length:
            spectrum = spectrum[:target_length]
        else:
            padding = target_length - len(spectrum)
            spectrum = np.pad(spectrum, (0, padding), 'constant')
            
    return spectrum

def calculate_cosine_similarity(vec1, vec2):
    """严格按照数学公式，使用基础运算计算两个向量的余弦相似性。"""
    numerator = np.sum(vec1 * vec2)
    denominator_part1 = np.sqrt(np.sum(vec1**2))
    denominator_part2 = np.sqrt(np.sum(vec2**2))
    denominator = denominator_part1 * denominator_part2
    if denominator == 0: return 0.0
    return numerator / denominator

def find_spectral_lines(spectrum, sensitivity):
    """
    在光谱中寻找吸收线（波谷）。
    sensitivity (prominence) 定义了谱线需要比周围突出多少才被识别。
    """
    # find_peaks 寻找波峰，所以我们将光谱数据取反来寻找波谷
    peaks, _ = find_peaks(-spectrum, prominence=sensitivity)
    return peaks

def compare_line_locations(lines1, lines2, tolerance):
    """
    比较两组谱线位置的相似度。
    使用 Dice 系数作为相似度分数: 2 * |A ∩ B| / (|A| + |B|)
    """
    if len(lines1) == 0 and len(lines2) == 0:
        return 1.0  # 两者都没有谱线，可以认为是完全匹配
    if len(lines1) == 0 or len(lines2) == 0:
        return 0.0  # 其中一个有谱线而另一个没有

    matches = 0
    # 为了避免重复匹配，创建一个已匹配谱线的标记
    matched_indices_in_lines2 = set()

    for l1 in lines1:
        for i, l2 in enumerate(lines2):
            if abs(l1 - l2) <= tolerance and i not in matched_indices_in_lines2:
                matches += 1
                matched_indices_in_lines2.add(i)
                break # 找到一个匹配后即跳出内层循环
                
    # 使用 Dice 系数计算相似度
    similarity_score = 2 * matches / (len(lines1) + len(lines2))
    return similarity_score


def analyze_and_display_spectra(image_path1, image_path2):
    """
    对两个恒星光谱进行分析，输出指定的图表和数值结果。
    """
    try:
        # --- 第1步: 加载并预处理光谱 ---
        spectrum1_raw = get_spectrum_from_image(image_path1)
        spectrum2_raw = get_spectrum_from_image(image_path2)
        target_len = max(len(spectrum1_raw), len(spectrum2_raw))
        spectrum1 = get_spectrum_from_image(image_path1, target_len)
        spectrum2 = get_spectrum_from_image(image_path2, target_len)
        norm_spec1 = normalize(spectrum1)
        norm_spec2 = normalize(spectrum2)

        # --- 第2步: 进行自卷积计算 ---
        autoconv1 = convolve(norm_spec1, norm_spec1, mode='same')
        autoconv2 = convolve(norm_spec2, norm_spec2, mode='same')
        norm_autoconv1 = normalize(autoconv1)
        norm_autoconv2 = normalize(autoconv2)

        # --- 第3步: 识别谱线并比较相似度 (新功能) ---
        lines1 = find_spectral_lines(norm_spec1, sensitivity=LINE_SENSITIVITY_THRESHOLD)
        lines2 = find_spectral_lines(norm_spec2, sensitivity=LINE_SENSITIVITY_THRESHOLD)
        line_location_similarity = compare_line_locations(lines1, lines2, tolerance=MATCHING_TOLERANCE_PIXELS)

        # --- 第4步: 计算两种余弦相似性 ---
        original_similarity = calculate_cosine_similarity(norm_spec1, norm_spec2)
        autoconv_similarity = calculate_cosine_similarity(norm_autoconv1, norm_autoconv2)

        # --- 第5步: 打印数值结果到终端 ---
        print("="*40)
        print("       恒星光谱相似度分析结果")
        print("="*40)
        print(f"设定参数:")
        print(f"  - 谱线识别灵敏度阈值: {LINE_SENSITIVITY_THRESHOLD}")
        print(f"  - 谱线匹配容差: {MATCHING_TOLERANCE_PIXELS} 像素\n")
        print(f"核心指标:")
        print(f"  - 谱线位置相似度: {line_location_similarity:.4f}  (检测到 {len(lines1)} vs {len(lines2)} 条谱线)")
        print(f"  - 原始光谱余弦相似度: {original_similarity:.6f}")
        print(f"  - 自卷积结果余弦相似度: {autoconv_similarity:.6f}")
        print("="*40)
        print("\n正在生成可视化图表...")

        # --- 第6步: 创建并显示可视化图表 ---
        fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(15, 12))
        fig.suptitle(image_path1.rsplit('.', 1)[0] + " VS " + image_path2.rsplit('.', 1)[0] \
                     + ' 恒星光谱分析', fontsize=16)

        # 图1: 原始光谱比较，并标记识别出的谱线
        ax1.plot(norm_spec1, label='光谱1 (原始)', color='royalblue')
        ax1.plot(norm_spec2, label='光谱2 (原始)', color='darkorange', linestyle='--')
        # 在图上用标记标出识别的谱线位置
        ax1.plot(lines1, norm_spec1[lines1], "x", color='red', markersize=8, label=f'光谱1谱线 (检出{len(lines1)}条)')
        ax1.plot(lines2, norm_spec2[lines2], "+", color='black', markersize=8, label=f'光谱2谱线 (检出{len(lines2)}条)')
        ax1.set_title(f'1. 原始光谱与谱线识别\n(曲线余弦相似度: {original_similarity:.4f})', fontsize=12)
        ax1.set_xlabel('像素位置 (代表波长)')
        ax1.set_ylabel('归一化亮度')
        ax1.legend()
        ax1.grid(True)

        # 图2: 自卷积结果比较
        ax2.plot(norm_autoconv1, label='光谱1 (自卷积)', color='forestgreen')
        ax2.plot(norm_autoconv2, label='光谱2 (自卷积)', color='crimson', linestyle='--')
        ax2.set_title(f'2. 自卷积结果比较\n(曲线余弦相似度: {autoconv_similarity:.4f})', fontsize=12)
        ax2.set_xlabel('像素位置')
        ax2.set_ylabel('归一化强度')
        ax2.legend()
        ax2.grid(True)
        
        plt.tight_layout(rect=[0, 0.03, 1, 0.95])
        plt.show()

    except FileNotFoundError as e:
        print(f"错误: 找不到文件 {e.filename}。请确保文件路径正确。")
    except Exception as e:
        print(f"处理过程中发生错误: {e}")

if __name__ == '__main__':
    image_file1 = 'B_Type_star_spectrogram_raw.png'
    image_file2 = 'Alkaid_pos.png'
    analyze_and_display_spectra(image_file1, image_file2)
