
(define *svm-max-epochs* 1001)		; default value
(define *svm-rate* 1.0)
(define *svm-max-alpha* #f)
(define *svm* #f)
(define *kernel-type* #f)
(define *svm-convergence* 0.01)

(define svm-svs first)
(define svm-bias second)

;;; Controls drawing
(define *svm-draw-interval-while-training* 500)

;; The actual names of the classes - will be converted to 0 and 1
(define *svm-class0* #f)		; set by svm-train
(define *svm-class1* #f)

(define (svm-train data kernel)
  (let ((classes (map car (class-counts data))))
    (cond ((= (length classes) 2)
	   (cond ((or (equal? classes '(0 1))
		      (equal? classes '(0.0 1.0))
		      (equal? classes '(1 0))
		      (equal? classes '(1.0 0.0)))
		  (set! *svm-class0* 0)
		  (set! *svm-class1* 1))
		 (else
		  (set! *svm-class0* (first classes))
		  (set! *svm-class1* (second classes)))))
	  (else 
	   (error "Only two classes allowed, but we have" classes))))
  (set! *kernel-type* kernel)
  (set! *svm*
	(kernel-adatron
	 *svm-max-epochs*
	 (svm-convert-training-data data *svm-class0* *svm-class1*)
	 kernel
	 *svm-rate*
	 *svm-max-alpha*))
  'done
  )

(define (svm-classify data-point)
  ;; This makes a prediction of the class, it needs to convert to
  ;; svm format and map the 0-1 prediction into the output classes.
  (let* ((converted 
	  (svm-convert-sample data-point *svm-class1*))
	 (out (make-svm-prediction (data-point-features converted)))
	 (prediction (if (>= out 0) *svm-class1* *svm-class0*)))
    (display* "The prediction is " prediction "(" out ")"
	      ".  Correct is " (data-point-class data-point)
	      "."
	      )
    prediction))

(define (make-svm-prediction point-features)
    (let ((sum 0))
      (for-each 
       (lambda (alpha.sv)
	 (set! sum (+ sum (* (data-point-class (cdr alpha.sv))
			     (car alpha.sv)
			     (kernel-value point-features
					   (data-point-features (cdr alpha.sv)))))))
       (svm-svs *svm*))
      (+ sum (svm-bias *svm*))))

(define (svm-convert-training-data data class-0 class-1)
  (random-reorder 
   (map (lambda (x) (svm-convert-sample x class-1))
	data)))

(define (svm-convert-sample x class-1)
  (cons (list (if (eq? (data-point-class x) class-1)
		  1.0 -1.0))
	(data-point-features x)))

(define (make-kernel size)
  (let ((kernel (make-vector size #f)))
    (do ((i 0 (+ i 1)))
	((= i size))
      (vector-set! kernel i (make-vector size 0.)))
    kernel))

(define (kernel-ref k i j)
  (vector-ref (vector-ref k i) j))

(define (kernel-set! k i j v)
  (vector-set! (vector-ref k i) j v))

(define (compute-kernel op data)
  (let* ((size (length data))
	 (k (make-kernel size)))
    (do ((i 0 (+ i 1))
	 (di data (cdr di)))
	((= i size))
      (do ((j i (+ j 1))
	   (dj di (cdr dj)))
	  ((= j size))
	(let ((v (op (car di) (car dj))))
	(kernel-set! k i j v)
	(kernel-set! k j i v)
	)))
    k)
  )

(define (dot v1 v2)
  (apply + (map * v1 v2)))

(define (compute-linear-kernel data)
  (compute-kernel (lambda (v1 v2) 
		    (dot (data-point-features v1) (data-point-features v2)))
		  data))

(define (poly v1 v2 d)
  (expt (+ 1 (dot v1 v2)) d))

(define (compute-polynomial-kernel data d)
  (compute-kernel (lambda (v1 v2) 
		    (poly (data-point-features v1) (data-point-features v2) d))
		  data))

(define (gauss v1 v2 sigma^2*2)
  (let ((d (map - v1 v2)))
    (exp (/ (- (dot d d)) sigma^2*2))))

(define (compute-gaussian-kernel data sigma)
  (let ((sigma^2*2 (* 2 sigma sigma)))
    (compute-kernel (lambda (v1 v2)
		      (gauss (data-point-features v1) (data-point-features v2)
			     sigma^2*2)) 
		    data)))

(define (kernel-value p1 p2)
  (cond ((eq? (first *kernel-type*) 'linear) 
	 (dot p1 p2))
	((eq? (first *kernel-type*) 'gaussian)
	 (gauss p1 p2 (second *kernel-type*)))
	((eq? (first *kernel-type*) 'polynomial)
	 (poly p1 p2 (second *kernel-type*)))
	(else (error "Unknown kernel type: " *kernel-type*))
	))

(define (vdot . vlist)
  (let ((size (vector-length (car vlist))))
    (do ((i 0 (+ i 1))
	 (sum 0.))
	((= i size) sum)
      (let ((prod 1))
	(for-each (lambda (v)
		    (set! prod (* prod (vector-ref v i))))
		  vlist)
	(set! sum (+ sum prod)))
      )))

(define (kernel-adatron-no-bias tmax data kernel-spec eta max-alpha)
  (let* ((size (length data))
	 (alpha (make-vector size 0))
	 (classes (list->vector (map data-point-class data)))
	 (kernel
	  (cond ((eq? (first kernel-spec) 'linear)
		 (compute-linear-kernel data))
		((eq? (first kernel-spec) 'gaussian)
		 (compute-gaussian-kernel data (second kernel-spec)))
		((eq? (first kernel-spec) 'polynomial)
		 (compute-polynomial-kernel data (second kernel-spec)))
		(else (error "Unknown kernel type: " kernel-spec))
		)))
    (do ((t 0 (+ t 1))
	 (margin 0))
	((or (= t tmax)
	     (< (abs (- 1 margin)) *svm-convergence*))
	 (set! *svm* (list (get-support-vectors-and-alpha alpha data) 0))
	 (if *draw* (draw-classifier svm-get-point data *window*))
	 (display* "Final margin=" margin)
	 *svm*) 
      (if (and *draw* *svm-draw-interval-while-training*)
	  (cond ((zero? (remainder t  *svm-draw-interval-while-training*))
		 (set! *svm* (list (get-support-vectors-and-alpha alpha data) 0))
		 (if (= t 0)
		     (draw-classifier svm-get-point data)
		     (draw-classifier svm-get-point data *window*)))))
      (do ((i 0 (+ i 1))
	   (z 0)
	   (zmin+ 100000)
	   (zmax- -100000))
	  ((= i size)
	   (set! margin (* 0.5 (- zmin+ zmax-)))
	   )
	(let ((class (vector-ref classes i))
	      (alp (vector-ref alpha i)))
	  (set! z (vdot alpha classes (vector-ref kernel i)))
	  ;; update zmin and zmax
	  (cond ((> class 0)
		 (if (and (< z zmin+) (if max-alpha (< alp max-alpha) #t)) 
		     (set! zmin+ z)))
		(else
		 (if (and (> z zmax-) (if max-alpha (< alp max-alpha) #t))
		     (set! zmax- z))))
	  (let* ((delta (* eta (- 1 (* z class))))
		 (new-alpha  (max 0 (+ alp delta))))
	    ;;(display* i ": alpha=" alp " delta=" delta " z=" z " zmin+=" zmin+ " zmax-=" zmax-) 
	    (vector-set! alpha i (if max-alpha (min max-alpha new-alpha) new-alpha)))
	  ))
      (if (= (remainder t 10) 0)
	  (display* "Iteration " t " ends with margin=" margin))
      )))

(define kernel-adatron kernel-adatron-no-bias)

(define (get-support-vectors-and-alpha alpha data)
  (let ((size (vector-length alpha)))
    (do ((i 0 (+ i 1))
	 (d data (cdr d))
	 (sva '()))
	((or (= i size) (null? data))
	 (reverse sva))
      (let ((a (vector-ref alpha i)))
	(if (not (= a 0))
	    (set! sva (cons (cons a (car d)) sva)))))))
	  