# Simeq_newton5.rb  solve nonlinear system of equations
#                   method: newton iteration using jacobian
#                   use nlin for higher order terms
# call   Simeq_newton5.new.simeq(n, nlin, a, y, var1, var2, var3, vari1, vari2, x)
#        neqn = n+nlim  a[neqn][neqn] * x[neqn] = y[neqn]  given a,y compute x
#
#  a maximal term can have up to power of 3 divided by power up to 2
#  in this implementation "5"
#  a term could be  x1 ,  x1*x3 ,  x2/x5 ,  x1*x2*x3/(x4*x5) ,  x2^3/x4^2 
#
#  Solve by initial guess at values of x1, x2, x3 computing products
#    x_next = x_initial - j_initial^-1 * (a *  x_initial - y)
#    in general x_next = x_prev - (j_prev^-1 * (a * x_prev - y))*b
#                                 where 0 < b < 1, often 0.5, for stability
#
#  solved when  abs sum each row a * x_next -y < epsilon
#
#  It may stall, stop if abs(x_next-x_prev)<epsilon
#  It may diverge, stop, indicate no solution
#                        (or try a different initial guess)
#  It may oscillate, stop, indicate no solution
#                        (or try a different initial guess)
#
#  the matrix equation could be:
#
#   x1    x2    x3    x1*x2  x3*x2^2  x2/x3  x1^3/x3^2 ... 
#   var1  var1  var1  var2   var3     varil  vari2
#   a1    a2    a3    a4     a5       a6     a7    coefficients for each equation
#   n=3 variables x1, x2, x3    
#   nlin=4 nonlinear terms  power 2, power 3, p0123/power 1  p0123/power 2
#   thus 3+4= 7 equations in 7 unknowns  (index 1 will be index 0 in code)
#
# | a1,1  a1,2  a1,3  a1,4   a1,5  a1,6  a1,7 ...| | x1         | |y1 |
# | a2,1  a2,2  a2,3  a2,4   a2,5  a2,6  a2,7 ...| | x2         | |y2 |
# | a3,1  a3,2  a3,3  a3,4   a3,5  a3,6  a3,7 ...| | x3         | |y3 |
# | a4,1  a4,2  a4,3  a4,4   a4,5  a4,6  a4,7 ...| | x1*x2      | |y4 |
# | a5,1  a5,2  a5,3  a5,4   a5,5  a5,6  a5,7 ...|*| x3*x2^2    |=|y5 |
# | a6,1  a6,2  a6,3  a6,4   a6,5  a6,6  a6,7 ...| | x2/x3      | |y6 |
# | a7,1  a7,2  a7,3  a7,4   a7,5  a7,6  a7,7 ...| | x1^3/x3^2  | |y7 |
# | ...                                       ...| | ...        | |...|
#       zero based variable numbers  x1 is 0, x2 is 1, x3 is 2, -1 means none
# var1    0   1   2   0   2   1   0
# var2   -1  -1  -1   1   1  -1   0
# var3   -1  -1  -1  -1   1  -1   0
# vari1  -1  -1  -1  -1  -1   2   2
# vari2-  1  -1  -1  -1  -1  -1   2
# 
#
#  the jacobian, j is the numerically computed derivatives based on a, x_prev
#  the derivative coefficients may be computed once based on a and the
#  unknown variables. the numeric value plugs in the x_prev values.
#
#  j1,1 = a1,1 + a1,4*x2 + 3*a1,7*x1^2/x3^2               deriv Row 1 wrt x1
#  j1,2 = a1,2 + a1,4*x1 + 2*a1,5*x3*x2 + a1,6/x3         deriv Row 1 wrt x2
#  j1,3 = a1,3 + a1,5*x2^2 + -a1,6*x2/x3^2 + -2*a1,7*x1^2/x3^3  
#                                                         deriv Row 1 wrt x3
#  j2,1 = a2,1 + a2,4*x2 + 3*a2,7*x1^2/x3^2               deriv Row 2 wrt x1
#  j2,2 = a2,2 + a2,4*x1 + 2*a2,5*x3*x2 + a2,6/x3         deriv Row 2 wrt x2
#  j2,3 = a1,3 + a2,5*x2^2 + -a2,6*x2/x3^2 + -2*a2,7*x1^2/x3^3  
#                                                         deriv Row 2 wrt x3
#  j3,1 = a3,1 + a3,4*x2 + 3*a3,7*x1^2/x3^2               deriv Row 3 wrt x1
#  j3,2 = a3,2 + a3,4*x1 + 2*a3,5*x3*x2 + a3,6/x3         deriv Row 3 wrt x2
#  j3,3 = a3,3 + a3,5*x2^2 + -a3,6*x2/x3^2 + -2*a3,7*x1^2/x3^3  
#                                                         deriv Row 3 wrt x3
#  ...
#  etc.
#
#  Code or a data structure must be available to know the equation
#  of the entries in the y vector. Symbolic computation of
#  derivatives is assumed to be available, when needed.
#

require_relative "Inverse"

class Simeq_newton5
  # derivatives computed from var1[], var2[], var3[], vari1[], vari2[] to
  # load ja, initial guess in x[0..n-1], returned solution in x[0..n+nlin-1]
  # n is number of variables, nlin is number of nonlinear terms
  # var1 etc are [0] to [n+nlin-1]

 
  def simeq(n, nlin, a, y, var1, var2, var3, vari1, vari2, x) # x[] solution returned
    eps = 1.0e-6 # (default)
    b = 1.0      # stability factor (default)
    b_init = 1.0 # user may set
    maxiter = 10  # maximum number of iterations (default)
    debug = 1    # should add monitor control
    varerr = 0   # input consistency flag
    # solve non linear simultaneous equations  a[][] * y[] = x[]
    # first n variables, then nlin non linear terms
    # var1, var2, var3, vari1, vari2 all length n+nlin  integer
    resid = 1.0e12  # residual from last iteration
    presid = 0.0    # residual from prior to last iteration
    ja = Array.new(n){Array.new(n)} # jacobian inverted in place
    jb = Array.new(n){Array.new(n)} # inverse
    x_resid = Array.new(n)
    next_resid = Array.new(n)
    x_tmp2 = Array.new(n+nlin)
    x_next = Array.new(n+nlin)
    t = 0.0 # term

    if debug>0
      puts "simeq_newton5 running n=#{n}, nlin=#{nlin}"
      for i in 0...n
        for j in 0...(n+nlin)
          puts "a[#{i}][#{j}]=#{a[i][j]}"
        end # j
        puts "y[#{i}]=#{y[i]}"
      end # i
    end # if
    b = b_init
    # check for consistent input
    for i in 0...n  # first n single variables -1 means unused
      if var1[i] != i
        puts "var1[#{i}]= #{var1[i]}  must be -1"
        varerr=1
      end # if
      if var2[i] != -1
        puts "var2[#{i}]= #{var2[i]}  must be -1"
        varerr=1
      end # if
      if var3[i] != -1
        puts "var3[#{i}]= #{var3[i]}  must be -1"
        varerr=1
      end # if
      if vari1[i] != -1
        puts "vari1[#{i}]= #{vari1[i]}  must be -1"
        varerr=1
      end # if
      if(vari2[i]!=-1)
        puts "vari2[#{i}]= #{vari2[i]}  must be -1"
        varerr=1
      end # if
    end # i single variable check
    for i in n...(n+nlin) # nlin nonlinear terms check
      if var1[i]<-1 || var1[i]>=n
        puts "var1[#{i}] must be -1 to #{(n-1)}"
        varerr=1
      end # if
      if var2[i]<-1 || var2[i]>=n
        puts "var2[#{i}] must be -1 to #{(n-1)}"
        varerr=1
      end # if
      if var3[i]<-1 || var3[i]>=n
        puts "var3[#{i}] must be -1 to #{(n-1)}"
        varerr=1
      end # if
      if vari1[i]<-1 || vari1[i]>=n
        puts "vari1[#{i}] must be -1 to #{(n-1)}"
        varerr=1
      end # if
      if vari2[i]<-1 || vari2[i]>=n
        puts "vari2["+i+"] must be -1 to #{(n-1)}"
        varerr=1
      end # if
    end # i nonlinear term check

    if varerr>0
      puts "var1  var2  var3  vari1 vari2"
      for i in 0...(n+nlin)
        puts "#{var1[i]} #{var2[i]} #{var3[i]} #{vari1[i]} #{vari2[i]}"
      end # i
      puts "simeq_newton5.rb  aborts, bad input"
      return
    end # if

    # setup nonlinear terms, this does not need to be set up by caller
    for i in n...(n+nlin)
      x[i] = 1.0 # there should be no entry for null term
      if var1[i]>=0
        x[i] = x[i] * x[var1[i]]
      end # if
      if var2[i]>=0
        x[i] = x[i] * x[var2[i]]
      end # if
      if var3[i]>=0
        x[i] = x[i] * x[var3[i]]
      end # if
      if vari1[i]>=0
        x[i] = x[i] / x[vari1[i]]
      end # if
      if vari2[i]>=0
        x[i] = x[i] / x[vari2[i]]
      end # if
    end # i
    # debug print
    if debug>0
      puts "initial guess at solution, can be bad"
      for i in 0...(n+nlin)
        puts "x[#{i}]=#{x[i]}"
      end # i
      puts " "
    end # if

    # compute residual
    resid = 0.0 
    for i in 0...n
      for j in 0...(n+nlin)
        x_resid[i] = a[i][j]*x[j] - y[i]
      end # j
      # x_resid used later if not solved
      resid = resid + (x_resid[i]).abs
    end # i
    # debug print
    if debug>0
      for i in 0...n
        puts "residual x_resid[#{i}]= #{x_resid[i]}"
      end # i
      puts " "
      puts "simeq_newton5 itr #{0},    initial residual=#{resid}"
    end # if
    presid = resid

    # iterate to find solution
    for itr in 1..maxiter
      # check for convergence
      if resid<eps
        break  # ???
      end # if

      # compute jacobian
      for i in 0...n # equation
	for j in 0...n  # variable
	  ja[i][j] = a[i][j] # each linear term
	  for k in n...(n+nlin)  # nonlinear term
            # add derivative of nonlinear term to ja[i][j] 
            # first group with no derivative of inverse term 
	    if vari1[k]!=j && vari2[k]!=j  # no derivative, may contribute
	      if var1[k]==j && var2[k]==j && var3[k]==j       # x^3  
	        t = a[i][k]*3.0*x[var2[k]]*x[var3[k]]         # 3x^2 
	        if vari1[k]>=0
                  t = t /x[vari1[k]]                          # vari1
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
	        ja[i][j] = ja[i][j] + t
	      elsif var1[k]==j && var2[k]==j                  # x^2 x 
	        t = a[i][k]*2.0*x[var2[k]]                    # 2x x  
	        if var3[k]>=0
                  t = t * x[var3[k]]                          # var3
                end # if
	        if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
	        ja[i][j] = ja[i][j] + t
	      elsif var2[k]==j && var3[k]==j                  # x x^2 
	        t = a[i][k]*2.0*x[var3[k]]                    # 2x x  
	        if var1[k]>=0
                  t = t * x[var1[k]]                          # var1
                end # if
	        if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
	        ja[i][j] = ja[i][j] + t
              elsif var3[k]==j && var1[k]==j                  # x x^2 
	        t = a[i][k]*2.0*x[var1[k]]                    # 2x x  
	        if var2[k]>=0
                  t = t * x[var2[k]]                          # var2
                end # if
	        if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
	        ja[i][j]  = ja[i][j] + t
              elsif var1[k]==j                                # only x 
		t = a[i][k]
	        if var2[k]>=0
                  t = t * x[var2[k]]                          # var2
                end # if
                if var3[k]>=0
                  t = t * x[var3[k]]                          # var3
                end # if
	        if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
		ja[i][j] = ja[i][j] +t
              elsif var2[k]==j                                # only x 
		t = a[i][k]
                if var1[k]>=0
                  t  = t * x[var1[k]]                         # var1
                end # if
                if var3[k]>=0
                  t = t * x[var3[k]]                          # var3
                end # if
	        if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
		ja[i][j] = ja[i][j] + t
              elsif var3[k]==j                                # only x 
		t = a[i][k]
	        if var2[k]>=0
                  t = t * x[var2[k]]                          # var2
                end # if
                if var1[k]>=0
                  t = t * x[var1[k]]                          # var1
                end # if  
	        if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if  
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
		ja[i][j] = ja[i][j] + t
	      end # if  no derivative of inverse term 
            elsif vari1[k]==j || vari2[k]==j # do derivative inverse term 
	      # three cases
	      # d/dx 1/x = -1/x^2    d/dx 1/x^2 = -2/x^3 
	      if vari1[k]==j && vari2[k]==j                   # 1/x^2 
		t = -2.0*a[i][k]/(x[vari1[k]]*x[vari1[k]]*x[vari1[k]])
		                                              # -2/x^3 
                if var1[k]>=0
                  t = t * x[var1[k]]                          # var1
                end # if
                if var2[k]>=0
                  t = t * x[var2[k]]                          # var2
                end # if
                if var3[k]>=0
                  t = t * x[var3[k]]                          # var3
                end # if
		ja[i][j] = ja[i][j] + t
	      elsif(vari1[k]==j)                              # 1/xi1 
		t = -a[i][k]/(x[vari1[k]]*x[vari1[k]])
		                                              # -1/x^2 
                if var1[k]>=0
                  t = t * x[var1[k]]                          # var1
                end # if
                if var2[k]>=0
                  t = t * x[var2[k]]                          # var2
                end # if
                if var3[k]>=0
                  t = t * x[var3[k]]                          # var3
                end # if
                if vari2[k]>=0
                  t = t / x[vari2[k]]                         # vari2
                end # if
		ja[i][j] = ja[i][j] + t
	      elsif vari2[k]==j                               # 1/xi2 
		t = -a[i][k]/(x[vari2[k]]*x[vari2[k]])        # -1/x^2 
                if var1[k]>=0
                  t = t * x[var1[k]]                          # var1
                end # if
                if var2[k]>=0
                  t = t * x[var2[k]]                          # var2
                end # if
                if var3[k]>=0
                  t = t * x[var3[k]]                          # var3
                end # if
                if vari1[k]>=0
                  t = t / x[vari1[k]]                         # vari1
                end # if
		ja[i][j] = ja[i][j] +t
	      end # if
	    end # if  of inverse derivative 
	  end # k
	end # j
      end # i
      if debug>0
        puts "ja computed "
        for i in 0...n
          for j in 0...n
	    puts "ja[#{i}][#{j}]=#{ja[i][j]}"
	  end # j
        end # i
        puts " "
      end # if

      # invert ja
      Inverse.new.invert(ja, jb)
      if debug>0
        puts "ja inverted "
        for i in 0...n
          for j in 0...n
            ja[i][j] = jb[i][j]
	    puts "ja[#{i}][#{j}]=#{ja[i][j]}"
	  end # j
        end # i
        puts " "
      end # if

      # x_resid = a x - y        error vector
      # x_tmp2 = j^-1 * x_resid  raw correction vector
      # x_next = x - (x_tmp2)*b weighted correction vector
      # x      = x_next         next x with computed higher terms

      for i in 0...n
	x_tmp2[i] = 0.0
        for j in 0...n
	  x_tmp2[i] = x_tmp2[i] + ja[i][j]*x_resid[j] 
        end # j
      end # i
      # debug print
      if debug>0
        for i in 0...n
	  puts "change x_tmp2[#{i}]=#{x_tmp2[i]}"
        end # i
        puts " "
      end # if
      for i in 0...n
        x_next[i] = x[i] - b*x_tmp2[i]
      end # i
      for i in n...(n+nlin)  # compute non linear terms
	x_next[i] = 1.0
	if var1[i]>=0
          x_next[i] = x_next[var1[i]]
        end # if
        if var2[i]>=0
          x_next[i] *= x_next[var2[i]]
        end # if
        if var3[i]>=0
          x_next[i] *= x_next[var3[i]]
        end # if
        if vari1[i]>=0
          x_next[i] /= x_next[vari1[i]]
        end # if
        if vari2[i]>=0
          x_next[i] /= x_next[vari2[i]]
        end # if
      end # i

      if debug>0
        for i in 0...(n+nlin)
          puts "x_next[#{i}]=#{x_next[i]}"
        end # i
        puts " "
      end # if

      # compute residual
      resid = 0.0 
      for i in 0...n
        next_resid[i] = 0.0
        for j in 0...(n+nlin)
	  next_resid[i] += a[i][j]*x_next[j]
        end # j
        next_resid[i] -= y[i] # x_resid used later if not solved
        resid += (next_resid[i]).abs
      end # i
      # debug print
      if debug>0
        for i in 0...n
	  puts "residual x_resid[#{i}]=#{x_resid[i]}"
        end # i
        puts " "
      end # if
      puts "simeq_newton5 itr #{itr}, prev=#{presid}, residual=#{resid}"
      if resid<presid # progress, increase b
	if b<1.0
          b = b * 1.4143
          if b>1.0
            b = 1.0
          end # if
          puts "b increased to #{b}"
	end # if
        # iterate
        presid = resid 
        for i in 0...(n+nlin)
          x[i] = x_next[i]
        end # i
        for i in 0...n
          x_resid[i] = next_resid[i]
        end # i
      else # worse, reduce b
	b = b*0.5
        puts "b reduced to #{b}"
      end # if
      if debug>0
        puts " "
      end # if
    end  # iteration    

    if debug>0
      if resid<eps
        puts "converged last residual = #{resid}"
      else
        puts "not converged last residual = #{resid}"
      end # if
      puts "Simeq_newton5.rb finished"
    end # if
  end # 
end # class simeq_newton5

