# -*- coding: utf-8 -*-
"""
Created on Fri Mar 19 12:34:03 2021

@author: jhoffman
"""

import sys
sys.path.append("/home/jhoffman/.local/bin")
import numpy as np
import xarray as xr
import matplotlib.pyplot as plt
import os
#os.system("pip install metpy --user")
#os.system("pip install netcdf4 --user")
import metpy.calc as mpcalc
from metpy.cbook import get_test_data
from metpy.units import units
from PIL import Image, ImageOps, ImageChops, ImageEnhance
from datetime import datetime, date
import glob
import random


xsize = 9024
ysize = 9024
xsize=xsize-2000+1
ysize=ysize-2000+1
blocksize=512

fn='/home/jhoffman/leads/2019005viirs.I1km.leads.nc'
print(fn)
data = xr.open_dataset(fn)
mask = data['Mask']
mask=mask[1000:xsize+2000-1000,1000:ysize+2000-1000]

im = Image.fromarray(np.array(mask == 100))
im.save("python.mask.png")
land=np.array([mask >= 200])
land=land*1
water=np.array([land == 0])
water=water*1
water=np.squeeze(water)

count=np.array([water == 2])
count=count*0.
count=np.squeeze(count)
for t in range(2,3):
  yy='2020'
  dd='01'
  c0=0
  print(t)
#  for m in range(1,5,2):
  for m in range(1,5):
    mm=m
    if m == 3 and t == 1: 
       continue
    if m != 3 and t == 2: 
       continue

    doy=date(int(yy),int(mm),int(dd)).timetuple().tm_yday
    mm=str(mm).zfill(2)
    dd=str(dd).zfill(2)
    doy=str(doy).zfill(3)
    yy=str(yy).zfill(4)
    print(yy, mm, dd)
    if t == 1: 
       tdata="/apollo/cloud/scratch/jhoffman/leads/training/Mcomp3/train"
    if t == 2: 
       tdata="/apollo/cloud/scratch/jhoffman/leads/test/Mcont3"
    fn= glob.glob('case4/'+yy+doy+'.*.MODIS.BT.png')
    print(fn)
    vfn= glob.glob(yy+doy+".test.vmask.png")

    ctestmax=np.amax(count)
    ctest=np.array(count)/ctestmax
    img = Image.fromarray((ctest * 255).astype(np.uint8))
    img=ImageOps.grayscale(img)
    ctestmax=ctestmax*100.
    ctestmax=str(int(ctestmax))

    img.save(yy+doy+ctestmax+"count.lead.test.png")
    
    for i in range(len(fn)):
      hh=fn[i].split('.')[-4]
      if hh == '24':
        continue
      print(fn[i], hh, c0)
      
      pic=Image.open(fn[i])
      BT=np.array(pic).astype(np.uint8)
      BT=(np.array(BT))*(np.array(BT < 254))

      im = Image.fromarray(BT)
      im=ImageOps.grayscale(im)
      enhancer =ImageEnhance.Contrast(im)
      im.save(yy+doy+"python.test.png")
      leadpic=Image.open(vfn[0])
      label=np.array(leadpic)
      lead=(np.array(label) < 200)
      lead=np.squeeze(lead)
      lead=lead*1*np.array(BT > 0)
      img = Image.fromarray((lead * 255).astype(np.uint8))
      img=ImageOps.grayscale(img)
      #img=ImageChops.invert(img)
      img.save(yy+doy+"python.lead.test.png")
      c=0
      cc=0
      ccc=0

      if np.amax(lead) == 0 :
        continue
      cmax=150
      if t == 1:
        cmax=1000
      elif t == 2:
        cmax=20
      else:
        cmax=20

      while c < cmax :
        ccc+=1
        if ccc > 1000:
           c = cmax +1

        istep0=int(random.uniform(0,xsize-blocksize))
        jstep0=int(random.uniform(0,ysize-blocksize))


        istep1=np.amin([istep0+blocksize,xsize-1])
        jstep1=np.amin([jstep0+blocksize,ysize-1])
        if istep1 >= xsize-1:
            istep0=istep1-blocksize
        if jstep1 >= ysize-1:
            jstep0=jstep1-blocksize
        iimg=np.empty((blocksize,blocksize))
        ilab=np.empty((blocksize,blocksize))
        ilab=lead[istep0:istep1,jstep0:jstep1]
        ilab[0:10,0:blocksize]=0
        ilab[blocksize-10:blocksize,0:blocksize]=0
        ilab[0:blocksize,0:10]=0
        ilab[0:blocksize,blocksize-10:blocksize]=0
        iimg=BT[istep0:istep1,jstep0:jstep1]
        if sum(sum(ilab)) < 5:
            continue

        if sum(sum(np.array(iimg > 0)*1)) < .15*blocksize*blocksize:
            continue
        if (np.amin(iimg) == np.amax(img)):
          continue
        img = Image.fromarray(iimg)
        img=ImageOps.grayscale(img)
        enhancer =ImageEnhance.Contrast(img)
        img = enhancer.enhance(3)
        count[istep0:istep1,jstep0:jstep1]=count[istep0:istep1,jstep0:jstep1]+0.01

        if t == 1:
          img.save(tdata+"/images/"+str(c0)+".png")
        elif t == 2:
          img.save(tdata+"/images/"+str(c0)+".png")
        else:
          img.save("data/test/images/"+str(c0)+".png")

        img = Image.fromarray((ilab*255).astype(np.uint8))
        img=ImageOps.grayscale(img)
        if t == 1:
            img.save(tdata+"/masks/"+str(c0)+".png")
        elif t == 2:
            img.save(tdata+"/masks/"+str(c0)+".png")
        else:
            img.save("data/test/truthmasks/"+str(c0)+".png")        

        c+=1
        c0+=1
        cc+=1 
ctestmax=np.amax(count)
ctest=np.array(count)/ctestmax
img = Image.fromarray((ctest * 255).astype(np.uint8))
img=ImageOps.grayscale(img)

ctestmax=ctestmax*100.
ctestmax=str(int(ctestmax))
print, ctestmax
img.save(ctestmax+"count.lead.test.png")
       
c=0
xi=int(xsize/(blocksize-25))
yi=int(ysize/(blocksize-25))
for ii in range(xi): 
    for jj in range(yi):
        istep0=(ii*blocksize)-(ii*25)
        jstep0=jj*blocksize-(jj*25)

        istep1=istep0+blocksize-1
        jstep1=jstep0+blocksize-1
        if istep1 > xsize-1:
            istep1=xsize-1
        if jstep1 > ysize-1:
            jstep1=ysize-1
        img=np.zeros((blocksize,blocksize))
        c += 1