Биорадар: Часть 2. Обучение ИНС реконструкции ЭКГ
作者
using MAT
using PyCall
using Plots
using FilePathsBase: mkpath, joinpath
using CSV, Tables
skimage = pyimport("skimage");
skt = pyimport("sklearn");
function load_data(filepath::String; offset::Int=0)
data = matread(filepath)
filename = basename(filepath)
radar_i = vec(data["radar_i"])[offset+1:end]
radar_q = vec(data["radar_q"])[offset+1:end]
fs_radar = data["fs_radar"]
ecg_data1 = data["tfm_ecg1"]
ecg_data2 = data["tfm_ecg2"]
fs_ecg = data["fs_ecg"]
return radar_i, radar_q, fs_radar, ecg_data1, ecg_data2, fs_ecg, filename
end
function EllipseReconstruction(i, q)
model = skimage.measure.EllipseModel();
stackData = hcat(i, q);
model.estimate(stackData)
xc, yc, a, b, theta = model.params;
IQ_meas = Complex.(i, q);
I = real.(IQ_meas) .- xc;
Q = imag.(IQ_meas) .- yc;
phi = -theta;
Arot = [[ cos(phi), sin(phi)],
[-sin(phi), cos(phi)]]
IQ1 = hcat(I, Q) * Arot;
IQ1 = mapreduce(permutedims, vcat, IQ1)
I2 = IQ1[:, 1] ./ a;
Q2 = IQ1[:, 2] ./ b;
circle = Complex.(I2, Q2);
I_comp = real.(circle);
Q_comp = imag.(circle);
sigma = angle.(circle);
lambda_ = 12.5e-3;
range_ = unwrap(sigma) ./ (4 * pi) * lambda_;
range_ .-= range_[1]
return I_comp, Q_comp, sigma, range_
end
function Filter(data, lowcut, highcut, fs, order)
responsetype = Bandpass(lowcut, highcut)
designmethod = Butterworth(order)
filter = digitalfilter(responsetype, designmethod; fs=fs)
return filtfilt(filter, data)
end
function getVitalSignals(radar_dist, fs_radar, order)
radar_respiration = Filter(radar_dist, 0.05, 0.5, fs_radar, order)
radar_pulse = Filter(radar_dist, 1.0, 6.0, fs_radar, order)
radar_heartsound = Filter(radar_dist, 10.0, 80.0, fs_radar, order)
radar_heartsound_denoise = skimage.restoration.denoise_wavelet(
radar_heartsound,
wavelet="db10",
mode="soft",
wavelet_levels=5,
method="VisuShrink",
rescale_sigma=true
)
return radar_respiration, radar_pulse, radar_heartsound, radar_heartsound_denoise
end
function plot_ecg_with_peaks(ecg_cleaned, waves, rpeaks; range_start=90000, range_end=100000)
# 1. Извлекаем подмножества
r_subset = filter(x -> range_start ≤ x ≤ range_end, rpeaks["ECG_R_Peaks"])
r_local = r_subset .- (range_start - 1)
p_subset = filter(x -> range_start ≤ x ≤ range_end, waves["ECG_P_Peaks"])
p_local = p_subset .- (range_start - 1)
q_subset = filter(x -> range_start ≤ x ≤ range_end, waves["ECG_Q_Peaks"])
q_local = q_subset .- (range_start - 1)
s_subset = filter(x -> range_start ≤ x ≤ range_end, waves["ECG_S_Peaks"])
s_local = s_subset .- (range_start - 1)
t_subset = filter(x -> range_start ≤ x ≤ range_end, waves["ECG_T_Peaks"])
t_local = t_subset .- (range_start - 1)
ecg_part = ecg_cleaned[range_start:range_end]
# 2. Рисуем график
p = plot(
layout = (1, 2),
size = (900, 400),
theme = :default,
framestyle = :box,
grid = true
)
# Левый: ECG + пики
plot!(p, ecg_part;
subplot = 1,
color = :black,
lw = 1.5,
label = "ECG",
xlabel = "Отсчёт",
ylabel = "Амплитуда",
title = "Фрагмент с пиками"
)
scatter!(p, r_local, ecg_cleaned[r_subset]; subplot = 1, marker = :circle, mc = :red, ms = 5, label = "R-пики")
scatter!(p, p_local, ecg_cleaned[p_subset]; subplot = 1, marker = :pentagon, mc = :green, ms = 6, label = "P-пики")
scatter!(p, q_local, ecg_cleaned[q_subset]; subplot = 1, marker = :diamond, mc = :blue, ms = 5, label = "Q-пики")
scatter!(p, s_local, ecg_cleaned[s_subset]; subplot = 1, marker = :utriangle,mc = :purple, ms = 5, label = "S-пики")
scatter!(p, t_local, ecg_cleaned[t_subset]; subplot = 1, marker = :dtriangle,mc = :orange, ms = 5, label = "T-пики")
# Правый: просто ECG
plot!(p, ecg_part;
subplot = 2,
color = :black,
lw = 1.2,
label = "",
xlabel = "Отсчёт",
title = "Чистый участок ЭКГ"
)
return p
end
const ϵ = 1f-8
function ncc_loss(ŷ::AbstractArray, y::AbstractArray)
μ̂ = mean(ŷ, dims=1)
μ = mean(y , dims=1)
ŷc = ŷ .- μ̂
yc = y .- μ
num = sum(ŷc .* yc, dims=1)
denom = sqrt.(sum(ŷc.^2, dims=1) .* sum(yc.^2, dims=1)) .+ ϵ
ncc = num ./ denom # –1 … +1
return mean(1 .- ncc) # 0 при полном совпадении
end
λ = 0.2f0 # вес штрафа – подберите
function loss(ŷ, y)
mse = Flux.mse(ŷ, y)
return mse + λ * ncc_loss(ŷ, y)
end
# function loss(y_pred::AbstractArray, y_true::AbstractArray)
# # α = 0.5
# sq = Flux.mse(y_pred, y_true)
# # sq = α*Flux.mse(y_pred, y_true) + (1-α)*Flux.huber_loss(y_pred, y_true)
# # mse_per_sample = mean(sq, dims=1)
# # err = sum(mse_per_sample)
# return sq
# end
function get_training_dir(dirpath::AbstractString, model_name::AbstractString)
model_dir = joinpath(dirpath, "model", model_name)
isdir(model_dir) || mkpath(model_dir)
existing = filter(name -> startswith(name, "train") &&
isdir(joinpath(model_dir, name)),
readdir(model_dir))
if !isempty(existing)
nums = parse.(Int, replace.(existing, "train" => ""))
newnum = maximum(nums) + 1
else
newnum = 1
end
name_new_dir = "train$newnum"
new_train_dir = joinpath(model_dir, name_new_dir)
mkpath(new_train_dir)
return new_train_dir, name_new_dir
end
function train!(
model, train_loader, opt, loss_fn,
device, epoch::Int, num_epochs::Int
)
running_loss = 0.0
n_batches = 0
for (i, (radar, ecg)) in enumerate(train_loader)
# 1) (W,C,B) → GPU
radar = reshape(radar, size(radar,1), 1, size(radar,2)) |> device
ecg = ecg |> device
# 2) loss + градиенты
loss_val, gs = Flux.withgradient(Flux.params(model)) do
ŷ = dropdims(model(radar); dims=2)
loss_fn(ŷ, ecg)
end
# 3) manual LR-scheduler по шагу
# 4) обновление оптимизатора (side-effect внутри)
Flux.update!(opt, Flux.params(model), gs)
running_loss += Float64(loss_val)
n_batches += 1
end
train_loss = running_loss / max(n_batches, 1)
return opt, train_loss
end
function mycollate(batch)
radar = hcat(first.(batch)...) # 5024 × N
ecg = hcat(last.(batch)...) # 5024 × N
return radar, ecg
end
function compute_metrics(pred, target; τ=0.05)
# 1) переводим на CPU и в плоский вектор Float64
p = vec(Array(pred))
t = vec(Array(target))
# 2) MAE
mae = skt.metrics.mean_absolute_error(t, p)
# 3) RMSE
rmse = sqrt(mean((p .- t).^2))
# 4) R² (коэффициент детерминации)
r2 = skt.metrics.r2_score(t, p)
# 5) точность: доля откликов с |error| ≤ τ
acc = mean(abs.(p .- t) .<= τ)
# 6) возвращаем результаты в NamedTuple
return (
MAE = mae,
RMSE = rmse,
R2 = r2,
Symbol("Acc@$(τ)") => acc
)
end
function validate(model, val_loader, loss_fn, device)
# Переключаем BatchNorm/Dropout в режим inference
Flux.testmode!(model)
running_loss = 0.0
preds = nothing
targets = nothing
n_batches = 0
for (radar, ecg) in val_loader
# 1) (W) → (W,C=1,B) → GPU
radar = reshape(radar, size(radar,1), 1, size(radar,2)) |> device
ecg = ecg |> device
# 2) Прямой проход (no gradient)
ŷ = dropdims(model(radar); dims=2) # (W, B)
# 3) Считаем лосс
loss_val = loss_fn(ŷ, ecg)
running_loss += Float64(loss_val)
n_batches += 1;
# 4) Сохраняем предсказания и цели
if preds === nothing
preds = ŷ
targets = ecg
else
preds = hcat(preds, ŷ) # конкатенируем по batch-оси
targets = hcat(targets, ecg)
end
Flux.trainmode!(model)
end
# 5) Считаем метрики
metrics = compute_metrics(preds, targets)
return running_loss / max(n_batches, 1) , metrics
end
function save_plots!(mae, rmse, r2, acc, train_loss, valid_loss;
model_dir::AbstractString,
run_name::AbstractString,
plot_name::AbstractString = "")
out_dir = joinpath(model_dir, plot_name, run_name, "outputs")
mkpath(out_dir)
# ─── маленькая вспом‑функция для одно‑линейных графиков ──────────
function line_plot(vec, ylabel, fname, col)
p = plot(vec; c=col, lw=2, xlabel="Epochs",
ylabel=ylabel, legend=false, size=(800,500))
savefig(p, joinpath(out_dir, fname))
end
line_plot(mae, "MAE", "$(plot_name)_MAE.png", :steelblue)
line_plot(rmse, "RMSE", "$(plot_name)_RMSE.png", :firebrick)
line_plot(r2, "R²", "$(plot_name)_R2.png", :seagreen)
line_plot(acc, "Acc", "$(plot_name)_Acc.png", :orange)
# ─── Loss‑ы на одном полотне ─────────────────────────────────────
p = plot(train_loss; c=:dodgerblue, label="train", lw=2,
xlabel="Epochs", ylabel="Loss", size=(800,500))
plot!(p, valid_loss; c=:crimson, label="val", lw=2)
savefig(p, joinpath(out_dir, "$(plot_name)_loss.png"))
metrics = Dict(
"MAE" => mae,
"RMSE" => rmse,
"R2" => r2,
"Acc" => acc,
"train_loss" => train_loss,
"valid_loss" => valid_loss,
)
# ─── CSV‑экспорт ────────────────────────────────────────────────
for (name, vec) in metrics
tbl = Tables.columntable((Epoch = 1:length(vec), Value = vec))
CSV.write(joinpath(out_dir, "$name.csv"), tbl)
end
println("✓ Пакет графиков и CSV сохранён в $(out_dir)")
end