# This Python 3 environment comes with many helpful analytics libraries installed
# It is defined by the kaggle/python docker image: https://github.com/kaggle/docker-python
# For example, here's several helpful packages to load in 

import numpy as np # linear algebra
import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)

# Input data files are available in the "../input/" directory.
# For example, running this (by clicking run or pressing Shift+Enter) will list the files in the input directory

import os
print(os.listdir("../input"))
# -*- coding:utf-8 -*-

import cv2
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
from scipy import ndimage
from skimage.measure import label,regionprops
import tensorflow as tf
import tensorflow.contrib.slim as slim
import os
print(os.getcwd())

flags = tf.flags
flags.DEFINE_string('DIR_HOME',r'D:\\kaggle\\airbus ship detect\\','')
flags.DEFINE_string('ckpt_dir',"results\\ckpt_dir\\",'model saved dir')
flags.DEFINE_string('savedimage_dir',"results\\predictedimage\\",'save predicted image')
flags.DEFINE_string('log_dir',r'E:\\program\\2.mechine learning\\kaggle\\Airbus_Ship_Detection\\results\\logs\\','save the output info')
FLAGS = flags.FLAGS

DIR_HOME = r'D:\kaggle\airbus ship detect\\'
TRAININGDATA_DIR = os.path.join(DIR_HOME,'train\\')
TESTDATA_DIR = os.path.join(DIR_HOME,'test\\')
SEGMENT_FILE = os.path.join(DIR_HOME,'train_ship_segmentations.csv')

##1.load dataset and process
boundaries = pd.read_csv(SEGMENT_FILE)
imagename_label = boundaries.copy()
print(boundaries.head(15))
#check ships
not_empty = pd.notna(boundaries['EncodedPixels'])

not_empty = not_empty.apply(lambda x: 1 if x==True else 0)
imagename_label['EncodedPixels'] = not_empty
imagename_label.to_csv('data\\imagename_label.csv')
# not_empty.hist()

def rle_decode(mask_rle,shape=(768,768)):
    s = mask_rle.split()
    starts,lengths = [np.asarray(x,dtype=int) for x in (s[0:][::2],s[1:][::2])]
    starts -= 1
    ends = starts +lengths
    img = np.zeros(shape[0]*shape[1],dtype=np.uint8)
    for lo,hi in zip(starts,ends):
        img[lo:hi] = 1
    return img.reshape(shape).T

def mask_as_image(image,in_mask_list):
    all_masks = np.zeros((768,768),dtype=np.int16)
    for mask in in_mask_list:
        if isinstance(mask,str):
            mask_decode = rle_decode(mask)
            all_masks += mask_decode
    return np.expand_dims(all_masks,-1)

# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode
def rle_encode(img, min_max_threshold=1e-3, max_mean_threshold=None):
    '''
    img: numpy array, 1 - mask, 0 - background
    Returns run length as string formated
    '''
    if np.max(img) < min_max_threshold:
        return 'None' ## no need to encode if it's all zeros
    if max_mean_threshold and np.mean(img) > max_mean_threshold:
        return '' ## ignore overfilled mask
    pixels = img.T.flatten()
    pixels = np.concatenate([[0], pixels, [0]])
    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1
    runs[1::2] -= runs[::2]
    return ' '.join(str(x) for x in runs)

def multi_rle_encode(img, **kwargs):
    '''
    Encode connected regions as separated masks
    '''
    labels = label(img)
    if img.ndim > 2:
        return [rle_encode(np.sum(labels==k, axis=2), **kwargs) for k in np.unique(labels[labels>0])]
    else:
        return [rle_encode(labels==k, **kwargs) for k in np.unique(labels[labels>0])]

def showSomeTrainImage():
    fig,axes = plt.subplots(3,5,figsize=(12,15))
    imagefiles = boundaries['ImageId']
    for i in range(3):
        for j in range(5):
            imagefile = imagefiles[i * 5 + j]
            image = cv2.imread(os.path.join(TRAININGDATA_DIR,imagefile))
            rle = boundaries.query('ImageId=="'+imagefile + '"')['EncodedPixels']
            mask = mask_as_image(image,rle)
            cv2.imwrite('hello.jpg',mask*255)

            lbl = label(mask)   #连通区域标记
            props = regionprops(lbl)
            image_1 = image.copy()
            image_1 = np.ones((768,768))
            for prop in props:
                cv2.rectangle(image_1,(prop.bbox[1],prop.bbox[0]),(prop.bbox[3],prop.bbox[2]),(0,255,255),2)
            axes[i][j].imshow(image_1)
            axes[i][j].set_xticks([])
            axes[i][j].set_title(imagefile)


def label_gen(imagefiles):
    label_images = []
    imagefiles = np.asarray(imagefiles,dtype=np.str)
    for imagefile in imagefiles:
        file = imagefile.rsplit("\\")[-1]
        image = cv2.imread(imagefile)
        rle = boundaries.query('ImageId=="'+ file + '"')['EncodedPixels']
        mask = mask_as_image(image,rle)
        label_images.append(np.reshape(mask,(768,768)))
    return np.asarray(label_images)
##data generator
def input_pipe(isTraining=True):
    filenames = []
    if isTraining:
        for root,subdir,files in os.walk(TRAININGDATA_DIR):
                filenames = [TRAININGDATA_DIR + filename for filename in files]
    else:
        for root,subdir,files in os.walk(TESTDATA_DIR):
                filenames = [TESTDATA_DIR + filename for filename in files]
    images_tensor = tf.convert_to_tensor(filenames,dtype=tf.string)
    input_queue = tf.train.slice_input_producer([images_tensor],num_epochs=None)

    images_content = tf.read_file(input_queue[0])
    images = tf.image.convert_image_dtype(tf.image.decode_png(images_content,channels=1),tf.float32)
    new_size = tf.constant([768,768],dtype=tf.int32)
    images = tf.image.resize_images(images,new_size)
    images_batch,label_batch = tf.train.shuffle_batch([images,input_queue[0]],batch_size = 8,capacity=500,min_after_dequeue=100)
    return images_batch,label_batch

def crop_and_concat(x1,x2):
    with tf.name_scope("crop_and_concat"):
        x1_shape = tf.shape(x1)
        x2_shape = tf.shape(x2)
        # offsets for the top left corner of the crop
        offsets = [0, (x1_shape[1] - x2_shape[1]) // 2, (x1_shape[2] - x2_shape[2]) // 2, 0]
        size = [-1, x2_shape[1], x2_shape[2], -1]
        x1_crop = tf.slice(x1, offsets, size)
        return tf.concat([x1_crop, x2], 3)

#find edge
def image_process(images_batch):
    from skimage.feature import canny
    images = []
    for image in images_batch:
        image = np.reshape(image,(768,768))
        # image = median(image)
        img = canny(image)
        images.append(np.reshape(img,(768,768,1)))
    return images

# selfconv to remove water ripple,but ship also being
def image_selfconv(image_batch):
    # image = image.convert('L')
    image_selfconv_batch = []
    for image in image_batch:
        cov_len = 30
        avg_img = []
        image = np.reshape(image,(768,768))
        for i in range(10):
            px = np.random.randint(0,256 - cov_len)
            py = np.random.randint(0,256 - cov_len)
            rect = (px,py,px+ cov_len,py+cov_len)
            # cut_img = image.crop(rect)
            cut_img = image[px:px+ cov_len,py:py+cov_len]
            avg_img.append(np.array(cut_img).reshape(-1))
        avg_img = np.array(avg_img).mean(axis=0).reshape(cov_len,cov_len)
        avg_img = avg_img /avg_img.sum()
        image = ndimage.convolve(image,avg_img)
        image_selfconv_batch.append(np.reshape(image,(768,768,1)))
    return np.asarray(image_selfconv_batch)


def Unet_v2(images):
    conv1_1 = slim.conv2d(images,8,[3,3],stride=1,padding='SAME',scope='conv1_1')
    conv1_1 = slim.batch_norm(conv1_1)

    conv1_2 = slim.conv2d(conv1_1,8,[3,3],stride=1,padding='SAME',scope='conv1_2')
    conv1_2 = slim.batch_norm(conv1_2)

    max_pool1 = slim.max_pool2d(conv1_2,[3,3],stride=2,padding='SAME',scope='maxpool1')

    conv2_1 = slim.conv2d(max_pool1,16,[3,3],stride=1,padding='SAME',scope='conv2_1')
    conv2_1 = slim.batch_norm(conv2_1)

    conv2_2 = slim.conv2d(conv2_1,16,[3,3],stride=1,padding='SAME',scope='conv2_2')
    conv2_2 = slim.batch_norm(conv2_2)

    max_pool2 = slim.max_pool2d(conv2_2,[3,3],stride=2,padding='SAME',scope='maxpool2')

    conv3_1 = slim.conv2d(max_pool2,32,[3,3],stride=1,padding='SAME',scope='conv3_1')
    conv3_1 = slim.batch_norm(conv3_1)

    conv3_2 = slim.conv2d(conv3_1,32,[3,3],stride=1,padding='SAME',scope='conv3_2')
    conv3_2 = slim.batch_norm(conv3_2)

    max_pool3 = slim.max_pool2d(conv3_2,[3,3],stride=2,padding='SAME',scope='maxpool3')

    conv4_1 = slim.conv2d(max_pool3,64,[3,3],stride=1,padding='SAME',scope='conv4_1')
    conv4_1 = slim.batch_norm(conv4_1)

    conv4_2 = slim.conv2d(conv4_1,64,[3,3],stride=1,padding='SAME',scope='conv4_2')
    conv4_2 = slim.batch_norm(conv4_2)

    max_pool4 = slim.max_pool2d(conv4_2,[3,3],stride=2,padding='SAME',scope='maxpool4')

    conv5_1 = slim.conv2d(max_pool4,128,[3,3],stride=1,padding='SAME',scope='conv5_1')
    conv5_1 = slim.batch_norm(conv5_1)

    conv5_2 = slim.conv2d(conv5_1,128,[3,3],stride=1,padding='SAME',scope='conv5_2')
    conv5_2 = slim.batch_norm(conv5_2)

    up_6_1 = slim.conv2d_transpose(conv5_2,128,[3,3],stride=2,padding='SAME',scope='up6_1')
    up6 = tf.concat((conv4_2,up_6_1),axis=-1)
    # up6 = crop_and_concat(conv4_2,up_6_1)
    conv6_1 = slim.conv2d(up6,64,[3,3],stride=1,padding='SAME',scope='conv6_1')
    conv6_1 = slim.batch_norm(conv6_1)

    conv6_2 = slim.conv2d(conv6_1,64,[3,3],stride=1,padding='SAME',scope='conv6_2')
    conv6_2 = slim.batch_norm(conv6_2)


    up_7_1 = slim.conv2d_transpose(conv6_2,64,[3,3],stride=2,padding='SAME',scope='up7_1')
    up7 = tf.concat((conv3_2,up_7_1),axis=-1)
    # up7 = crop_and_concat(conv3_2,up_7_1)
    conv7_1 = slim.conv2d(up7,32,[3,3],stride=1,padding='SAME',scope='conv7_1')
    conv7_1 = slim.batch_norm(conv7_1)

    conv7_2 = slim.conv2d(conv7_1,32,[3,3],stride=1,padding='SAME',scope='conv7_2')
    conv7_2 = slim.batch_norm(conv7_2)


    up_8_1 = slim.conv2d_transpose(conv7_2,32,[3,3],stride=2,padding='SAME',scope='up8_1')
    up8 = tf.concat((conv2_2,up_8_1),axis=-1)
    # up8 = crop_and_concat(conv2_2,up_8_1)
    conv8_1 = slim.conv2d(up8,16,[3,3],stride=1,padding='SAME',scope='conv8_1')
    conv8_1 = slim.batch_norm(conv8_1)

    conv8_2 = slim.conv2d(conv8_1,16,[3,3],stride=1,padding='SAME',scope='conv8_2')
    conv8_2 = slim.batch_norm(conv8_2)

    up_9_1 = slim.conv2d_transpose(conv8_2,16,[3,3],stride=2,padding='SAME',scope='up9_1')
    up9 = tf.concat((conv1_2,up_9_1),axis=-1)
    # up9 = crop_and_concat(conv1_2,up_9_1)
    conv9_1 = slim.conv2d(up9,8,[3,3],stride=1,padding='SAME',scope='conv9_1')
    conv9_1 = slim.batch_norm(conv9_1)

    conv9_2 = slim.conv2d(conv9_1,8,[3,3],stride=1,padding='SAME',scope='conv9_2')
    conv9_2 = slim.batch_norm(conv9_2)

    predicts = slim.conv2d(conv9_2,2,[1,1],stride=1,padding='SAME',scope='conv9_3') #最后一维输出为2,class num,从每个像素上来分类：是否为船体
    return predicts

def Unet(images):
    comp0 = slim.avg_pool2d(images,(6,6),stride=6,padding='SAME')

    conv1_1 = slim.conv2d(comp0,16,[3,3],stride=1,padding='SAME',scope='conv1_1')
    conv1_1 = slim.batch_norm(conv1_1)

    conv1_2 = slim.repeat(conv1_1,1,slim.conv2d,16,[3,3],scope='conv1_2')
    conv1_2 = slim.batch_norm(conv1_2)

    max_pool1 = slim.max_pool2d(conv1_2,[3,3],stride=2,padding='SAME',scope='maxpool1')

    conv2_1 = slim.conv2d(max_pool1,32,[3,3],stride=1,padding='SAME',scope='conv2_1')
    conv2_1 = slim.batch_norm(conv2_1)

    conv2_2 = slim.repeat(conv2_1,1,slim.conv2d,32,[3,3],scope='conv2_2')
    conv2_2 = slim.batch_norm(conv2_2)

    max_pool2 = slim.max_pool2d(conv2_2,[3,3],stride=2,padding='SAME',scope='maxpool2')

    conv3_1 = slim.conv2d(max_pool2,64,[3,3],stride=1,padding='SAME',scope='conv3_1')
    conv3_1 = slim.batch_norm(conv3_1)

    conv3_2 = slim.repeat(conv3_1,1,slim.conv2d,64,[3,3],scope='conv3_2')
    conv3_2 = slim.batch_norm(conv3_2)

    max_pool3 = slim.max_pool2d(conv3_2,[3,3],stride=2,padding='SAME',scope='maxpool3')

    conv4_1 = slim.conv2d(max_pool3,128,[3,3],stride=1,padding='SAME',scope='conv4_1')
    conv4_1 = slim.batch_norm(conv4_1)

    conv4_2 = slim.repeat(conv4_1,1,slim.conv2d,128,[3,3],scope='conv4_2')
    conv4_2 = slim.batch_norm(conv4_2)

    max_pool4 = slim.max_pool2d(conv4_2,[3,3],stride=2,padding='SAME',scope='maxpool4')

    conv5_1 = slim.conv2d(max_pool4,256,[3,3],stride=1,padding='SAME',scope='conv5_1')
    conv5_1 = slim.batch_norm(conv5_1)

    conv5_2 = slim.repeat(conv5_1,1,slim.conv2d,256,[3,3],scope='conv5_2')
    conv5_2 = slim.batch_norm(conv5_2)

    up_6_1 = slim.conv2d_transpose(conv5_2,256,[3,3],stride=2,padding='SAME',scope='up6_1')
    up_6_2 = slim.conv2d(up_6_1,128,[3,3],stride=1,padding='SAME',scope='up6_2')
    up_6_2 = slim.batch_norm(up_6_2)
    up6 = tf.concat((conv4_2,up_6_2),axis=-1)
    # up6 = crop_and_concat(conv4_2,up_6_1)
    conv6_1 = slim.conv2d(up6,128,[3,3],stride=1,padding='SAME',scope='conv6_1')
    conv6_1 = slim.batch_norm(conv6_1)
    conv6_2 = slim.repeat(conv6_1,1,slim.conv2d,128,[3,3],scope='conv6_2')
    conv6_2 = slim.batch_norm(conv6_2)

    up_7_1 = slim.conv2d_transpose(conv6_2,128,[3,3],stride=2,padding='SAME',scope='up7_1')
    up_7_2 = slim.conv2d(up_7_1,64,[3,3],stride=1,padding='SAME',scope='up7_2')
    up_7_2 = slim.batch_norm(up_7_2)
    up7 = tf.concat((conv3_2,up_7_2),axis=-1)
    # up7 = crop_and_concat(conv3_2,up_7_1)

    conv7_1 = slim.conv2d(up7,64,[3,3],stride=1,padding='SAME',scope='conv7_1')
    conv7_1 = slim.batch_norm(conv7_1)
    conv7_2 = slim.repeat(conv7_1,1,slim.conv2d,64,[3,3],scope='conv7_2')
    conv7_2 = slim.batch_norm(conv7_2)


    up_8_1 = slim.conv2d_transpose(conv7_2,64,[3,3],stride=2,padding='SAME',scope='up8_1')
    up_8_2 = slim.conv2d(up_8_1,32,[3,3],stride=1,padding='SAME',scope='up8_2')
    up_8_2 = slim.batch_norm(up_8_2)
    up8 = tf.concat((conv2_2,up_8_2),axis=-1)
    # up8 = crop_and_concat(conv2_2,up_8_1)
    conv8_1 = slim.conv2d(up8,32,[3,3],stride=1,padding='SAME',scope='conv8_1')
    conv8_1 = slim.batch_norm(conv8_1)

    conv8_2 = slim.repeat(conv8_1,1,slim.conv2d,32,[3,3],scope='conv8_2')
    conv8_2 = slim.batch_norm(conv8_2)

    up_9_1 = slim.conv2d_transpose(conv8_2,32,[3,3],stride=2,padding='SAME',scope='up9_1')
    up_9_2 = slim.conv2d(up_9_1,16,[3,3],stride=1,padding='SAME',scope='up9_2')
    up_9_2 = slim.batch_norm(up_9_2)
    up9 = tf.concat((conv1_2,up_9_2),axis=-1)
    # up9 = crop_and_concat(conv1_2,up_9_1)
    conv9_1 = slim.conv2d(up9,16,[3,3],stride=1,padding='SAME',scope='conv9_1')
    conv9_1 = slim.batch_norm(conv9_1)

    conv9_2 = slim.repeat(conv9_1,1,slim.conv2d,16,[3,3],scope='conv9_2')
    conv9_2 = slim.batch_norm(conv9_2)

    dcmp10 = slim.conv2d_transpose(conv9_2,16,[3,3],stride=6,scope='dcmp10')
    mrge10 = tf.concat((images,dcmp10),axis=-1)

    conv10_1 = slim.conv2d(mrge10,16,[3,3],stride=1,padding='SAME',scope='conv10_1')
    conv10_2 = slim.conv2d(conv10_1,8,[3,3],stride=1,padding='SAME',scope='conv10_2')

    predicts = slim.conv2d(conv10_2,2,[1,1],activation_fn=tf.nn.sigmoid,stride=1,padding='SAME',scope='conv11') #最后一维输出为2,class num,从每个像素上来分类：是否为船体
    return predicts

def loss_crossentropy(predicts,labels):
    # loss1 = tf.nn.sparse_softmax_cross_entropy_with_logits(logits=predicts,labels=labels)
    loss1 = tf.nn.sparse_softmax_cross_entropy_with_logits(logits=predicts,labels=labels)
    # loss1 = tf.pow(tf.cast(labels,dtype=tf.float32),predicts)
    # loss2 = Iou(labels,predicts)
    mean_loss = tf.reduce_mean(loss1)

    correction_pred = tf.equal(tf.argmax(input=predicts,axis=3,output_type=tf.int32),labels)
    accuracy = tf.reduce_mean(tf.cast(correction_pred,dtype=tf.float32,name='accuracy'))

    global_step = tf.get_variable("step", [], initializer=tf.constant_initializer(0), trainable=False)
    rate = tf.train.exponential_decay(2e-3, global_step, decay_steps=100, decay_rate=0.95, staircase=True)
    optimizer = tf.train.AdamOptimizer(learning_rate=rate).minimize(mean_loss)

    return mean_loss,optimizer,accuracy

def Iou(y_true,y_pred,eps = 1e-6):
    y_true_f = slim.flatten(y_true)
    y_pred_f = slim.flatten(y_pred)
    intersection = y_true_f * y_pred_f
    union = tf.reduce_sum(y_true_f) + tf.reduce_sum(y_pred_f)
    return (2.0 * intersection + eps) / ( union+ eps)

def train():
    with tf.Session() as sess:
        images = tf.placeholder(dtype= tf.float32,shape=[None,768,768,1],name='input')
        labels = tf.placeholder(dtype=tf.int32,shape=[None,768,768],name='label')

        images_batch_tensor,label_batch_tensor = input_pipe()
        predict_tensor = Unet(images)
        loss_tensor,optimizer_tensor,accuracy_tensor = loss_crossentropy(predict_tensor,labels)
        tf.global_variables_initializer().run()
        coord = tf.train.Coordinator()
        threads = tf.train.start_queue_runners(sess=sess,coord=coord)
        saver = tf.train.Saver(max_to_keep=3)
        model_file = tf.train.latest_checkpoint(FLAGS.ckpt_dir)
        step = 0
        if model_file:
            saver.restore(sess,model_file)
            step = int(model_file.split('-')[1])

        print("=============== begin training.===============")
        try:
            while not coord.should_stop():
                train_batch_image,train_batch_files = sess.run([images_batch_tensor,label_batch_tensor])
                train_batch_image = image_selfconv(train_batch_image)
                train_batch_label = label_gen(train_batch_files)
                for label_mask in train_batch_label:
                    multi_rle_encode(label_mask)
                _,loss,predicts,accuracy = sess.run([optimizer_tensor,loss_tensor,predict_tensor,accuracy_tensor],feed_dict={
                    images:train_batch_image,
                    labels:train_batch_label
                })
                for masks,pred_label,file in zip(train_batch_label,predicts,np.asarray(train_batch_files,dtype=np.str)):
                    filename = file.split("\\")[-1]
                    # np.save("results\\predictedimage\\"+filename.split('.')[0] + '_mask',masks)
                    # np.save("results\\predictedimage\\"+filename.split('.')[0] + '_pred',pred_label)
                    # cv2.imwrite("results\\predictedimage\\"+filename.split('.')[0] + '_mask.jpg',masks[:,:] * 255)
                    # cv2.imwrite("results\\predictedimage\\"+filename.split('.')[0] + '_pred.jpg',pred_label[:,:]* 255)
                step += 1
                stroutput = 'training step: %d, loss: %f, training accuracy:%f'\
                            % (step,loss,accuracy)
                print(stroutput)
                with open(FLAGS.log_dir + 'log.txt','a') as pf:
                    pf.write(stroutput + '\n')
                if step % 5 == 0:
                    saver.save(sess,FLAGS.ckpt_dir + 'model.ckpt',global_step=step)
        except tf.errors.OutOfRangeError:
            print("==================> training fanished.")
        finally:
            coord.request_stop()
        coord.join(threads)

def predict():
    pred_rows = []
    with tf.Session() as sess:
        images = tf.placeholder(dtype= tf.float32,shape=[None,768,768,3],name='input')
        images_tensor,image_name_tensor = input_pipe(isTraining=False)
        predicts = Unet(images)
        tf.global_variables_initializer().run()
        saver = tf.train.Saver()
        model_file = tf.train.latest_checkpoint(FLAGS.ckpt_dir)
        coord = tf.train.Coordinator()
        threads = tf.train.start_queue_runners(sess=sess,coord=coord)

        if model_file:
            saver.restore(sess,model_file)
            print("=============== begin test.===============")
            try:
                while not coord.should_stop():
                    test_images_batch,test_images_name_batch = sess.run([images_tensor,image_name_tensor])
                    images_pred = sess.run(tf.argmax(predicts,axis=3),feed_dict = {
                        images:test_images_batch
                    })
                    for image,image_name in zip(images_pred,test_images_name_batch):
                        image_name = image_name.decode()
                        savedimagename = image_name.split('\\')[-1].split('.')[0]
                        cv2.imwrite(FLAGS.savedimage_dir + savedimagename + '.jpg',image * 255)

                        rles = multi_rle_encode(image)
                        if len(rles)>0:
                            for rle in rles:
                                pred_rows += [{'ImageId': savedimagename + '.jpg', 'EncodedPixels': rle}]
                            else:
                                pred_rows += [{'ImageId': savedimagename + '.jpg', 'EncodedPixels': None}]

            except tf.errors.OutOfRangeError:
                print("==================> test fanished.")
                submission_df = pd.DataFrame(pred_rows)[['ImageId', 'EncodedPixels']]
                submission_df.to_csv('submission.csv', index=False)

            finally:
                coord.request_stop()
            coord.join(threads)

if __name__ == "__main__":
    train()
    predict()

# Any results you write to the current directory are saved as output.