#!/usr/bin/python
import cereal
from cereal import *

def add(cereal,other):
	if not is_cereal(cereal):
		print "you need a cereal array"
		return
	if is_cereal(other):
		if (get_width(cereal) > get_width(other)):
			width = get_width(other)
		else:
			width = get_width(cereal)
		if (get_height(cereal) > get_height(other)):
			height = get_height(other)
		else:
			height = get_height(cereal)
		if get_nchannel(cereal) > get_nchannel(other):
			length = get_nchannel(other)
		else:
			length = get_nchannel(cereal)
		new_cereal = []
		for y in range(height):
			new_row = []
			for x in range(width):
				new_value = ()
				for n in range(length):
					new_value += (cereal[y][x][n] + other[y][x][n],)
				new_row.append(new_value)
			new_cereal.append(new_row)
		return new_cereal
	elif (type(other) == float):
		new_cereal = []
		for y in range(len(cereal)):
			new_row = []
			for x in range(len(cereal[y])):
				new_value = ()
				for n in range(len(cereal[y][x])):
					new_value += (cereal[y][x][n] + other,)
				new_row.append(new_value)
			new_cereal.append(new_row)
		return new_cereal
	else:
		print "the other must be a float or an array"
	return

def mul(cereal,other):
	if not is_cereal(cereal):
		print "you need a cereal array"
		return
	if is_cereal(other):
		if (get_width(cereal) > get_width(other)):
			width = get_width(other)
		else:
			width = get_width(cereal)
		if (get_height(cereal) > get_height(other)):
			height = get_height(other)
		else:
			height = get_height(cereal)
		if get_nchannel(cereal) > get_nchannel(other):
			length = get_nchannel(other)
		else:
			length = get_nchannel(cereal)
		new_cereal = []
		for y in range(height):
			new_row = []
			for x in range(width):
				new_value = ()
				for n in range(length):
					new_value += (cereal[y][x][n] * other[y][x][n],)
				new_row.append(new_value)
			new_cereal.append(new_row)
		return new_cereal
	elif (type(other) == float):
		new_cereal = []
		for y in range(len(cereal)):
			new_row = []
			for x in range(len(cereal[y])):
				new_value = ()
				for n in range(len(cereal[y][x])):
					new_value += (cereal[y][x][n] * other,)
				new_row.append(new_value)
			new_cereal.append(new_row)
		return new_cereal
	else:
		print "the other must be a float or an array"
	return

def div(cereal,other):
	if not is_cereal(cereal):
		print "you need a cereal array"
		return
	if is_cereal(other):
		if (get_width(cereal) > get_width(other)):
			width = get_width(other)
		else:
			width = get_width(cereal)
		if (get_height(cereal) > get_height(other)):
			height = get_height(other)
		else:
			height = get_height(cereal)
		if get_nchannel(cereal) > get_nchannel(other):
			length = get_nchannel(other)
		else:
			length = get_nchannel(cereal)
		new_cereal = []
		for y in range(height):
			new_row = []
			for x in range(width):
				new_value = ()
				for n in range(length):
					new_value += (cereal[y][x][n] / other[y][x][n],)
				new_row.append(new_value)
			new_cereal.append(new_row)
		return new_cereal
	elif (type(other) == float):
		new_cereal = []
		for y in range(len(cereal)):
			new_row = []
			for x in range(len(cereal[y])):
				new_value = ()
				for n in range(len(cereal[y][x])):
					new_value += (cereal[y][x][n] / other,)
				new_row.append(new_value)
			new_cereal.append(new_row)
		return new_cereal
	else:
		print "the other must be a float or an array"
	return

if __name__ == "__main__":
	pass
