import hyperspy.api as hs
import atomap.api as am
import pixstem.api as ps

# Load Medipix3 data
s = ps.load_ps_signal("b002_011_subset_corrected.hspy", lazy=True)

# Load atom_lattice for the atom positions
atom_lattice = am.load_atom_lattice_from_hdf5("b005_011_atom_lattice.hdf5", construct_zone_axes=False)
sublattice_A = atom_lattice.sublattice_list[0]
sublattice_B = atom_lattice.sublattice_list[1]
sublattice_O = atom_lattice.sublattice_list[2]

sublattice_A.construct_zone_axes(atom_plane_tolerance=0.7)
sublattice_B.construct_zone_axes(atom_plane_tolerance=0.7)
sublattice_O.construct_zone_axes(atom_plane_tolerance=0.7)

################ A-cations
zone_vector_111_A = sublattice_A.zones_axis_average_distances[3]

# Get STO A cation
plane_111_A_STO = sublattice_A.atom_planes_by_zone_vector[zone_vector_111_A][3]
A_STO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_A_STO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    A_STO_111_diff_list.append(s_diff)
s_A_STO_111_diff_stack = hs.stack(A_STO_111_diff_list)
s_A_STO_111_diff = s_A_STO_111_diff_stack.mean(0)
s_A_STO_111_diff.metadata.xlist = x_list
s_A_STO_111_diff.metadata.ylist = y_list

# Get LFO A cation
plane_111_A_LFO = sublattice_A.atom_planes_by_zone_vector[zone_vector_111_A][23]
A_LFO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_A_LFO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    A_LFO_111_diff_list.append(s_diff)
s_A_LFO_111_diff_stack = hs.stack(A_LFO_111_diff_list)
s_A_LFO_111_diff = s_A_LFO_111_diff_stack.mean(0)
s_A_LFO_111_diff.metadata.xlist = x_list
s_A_LFO_111_diff.metadata.ylist = y_list

# Get LSMO A cation
plane_111_A_LMO = sublattice_A.atom_planes_by_zone_vector[zone_vector_111_A][43]
A_LMO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_A_LMO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    A_LMO_111_diff_list.append(s_diff)
s_A_LMO_111_diff_stack = hs.stack(A_LMO_111_diff_list)
s_A_LMO_111_diff = s_A_LMO_111_diff_stack.mean(0)
s_A_LMO_111_diff.metadata.xlist = x_list
s_A_LMO_111_diff.metadata.ylist = y_list

################ B-cations
zone_vector_111_B = sublattice_B.zones_axis_average_distances[3]

# Get STO B cation
plane_111_B_STO = sublattice_B.atom_planes_by_zone_vector[zone_vector_111_B][2]
B_STO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_B_STO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    B_STO_111_diff_list.append(s_diff)
s_B_STO_111_diff_stack = hs.stack(B_STO_111_diff_list)
s_B_STO_111_diff = s_B_STO_111_diff_stack.mean(0)
s_B_STO_111_diff.metadata.xlist = x_list
s_B_STO_111_diff.metadata.ylist = y_list

# Get LFO B cation
plane_111_B_LFO = sublattice_B.atom_planes_by_zone_vector[zone_vector_111_B][21]
B_LFO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_B_LFO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    B_LFO_111_diff_list.append(s_diff)
s_B_LFO_111_diff_stack = hs.stack(B_LFO_111_diff_list)
s_B_LFO_111_diff = s_B_LFO_111_diff_stack.mean(0)
s_B_LFO_111_diff.metadata.xlist = x_list
s_B_LFO_111_diff.metadata.ylist = y_list

# Get LSMO B cation
plane_111_B_LMO = sublattice_B.atom_planes_by_zone_vector[zone_vector_111_B][41]
B_LMO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_B_LMO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    B_LMO_111_diff_list.append(s_diff)
s_B_LMO_111_diff_stack = hs.stack(B_LMO_111_diff_list)
s_B_LMO_111_diff = s_B_LMO_111_diff_stack.mean(0)
s_B_LMO_111_diff.metadata.xlist = x_list
s_B_LMO_111_diff.metadata.ylist = y_list

################ Oxygen
zone_vector_111_O = sublattice_O.zones_axis_average_distances[3]

# Get STO Oxygen
plane_111_O_STO = sublattice_O.atom_planes_by_zone_vector[zone_vector_111_O][1]
O_STO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_O_STO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    O_STO_111_diff_list.append(s_diff)
s_O_STO_111_diff_stack = hs.stack(O_STO_111_diff_list)
s_O_STO_111_diff = s_O_STO_111_diff_stack.mean(0)
s_O_STO_111_diff.metadata.xlist = x_list
s_O_STO_111_diff.metadata.ylist = y_list

# Get LFO Oxygen
plane_111_O_LFO = sublattice_O.atom_planes_by_zone_vector[zone_vector_111_O][21]
O_LFO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_O_LFO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    O_LFO_111_diff_list.append(s_diff)
s_O_LFO_111_diff_stack = hs.stack(O_LFO_111_diff_list)
s_O_LFO_111_diff = s_O_LFO_111_diff_stack.mean(0)
s_O_LFO_111_diff.metadata.xlist = x_list
s_O_LFO_111_diff.metadata.ylist = y_list

# Get LMO Oxygen
plane_111_O_LMO = sublattice_O.atom_planes_by_zone_vector[zone_vector_111_O][41]
O_LMO_111_diff_list = []
x_list, y_list = [], []
for atom in plane_111_O_LMO.atom_list:
    x = int(round(atom.pixel_x))
    y = int(round(atom.pixel_y))
    x_list.append(x)
    y_list.append(y)
for x, y in zip(x_list, y_list):
    s_diff = s.inav[x, y]
    s_diff.compute(progressbar=False)
    O_LMO_111_diff_list.append(s_diff)
s_O_LMO_111_diff_stack = hs.stack(O_LMO_111_diff_list)
s_O_LMO_111_diff = s_O_LMO_111_diff_stack.mean(0)
s_O_LMO_111_diff.metadata.xlist = x_list
s_O_LMO_111_diff.metadata.ylist = y_list

################## Save signals
s_A_STO_111_diff.save("b006_011_A_STO_diff.hspy", overwrite=True)
s_B_STO_111_diff.save("b006_011_B_STO_diff.hspy", overwrite=True)
s_O_STO_111_diff.save("b006_011_O_STO_diff.hspy", overwrite=True)

s_A_LFO_111_diff.save("b006_011_A_LFO_diff.hspy", overwrite=True)
s_B_LFO_111_diff.save("b006_011_B_LFO_diff.hspy", overwrite=True)
s_O_LFO_111_diff.save("b006_011_O_LFO_diff.hspy", overwrite=True)

s_A_LMO_111_diff.save("b006_011_A_LMO_diff.hspy", overwrite=True)
s_B_LMO_111_diff.save("b006_011_B_LMO_diff.hspy", overwrite=True)
s_O_LMO_111_diff.save("b006_011_O_LMO_diff.hspy", overwrite=True)
