#include <stdio.h>
#include <math.h>

#define N 4

void pivot(double x[][N+1], int i) {
    double c[N+1] = {0};
    for (int j = i + 1; j < N; j++) {
        if (fabs(x[i][i]) < fabs(x[j][i])) {
            for (int n = 0; n < N + 1; n++) {
                c[n] = x[j][n];
                x[j][n] = x[i][n];
                x[i][n] = c[n];
            }
        }
    }
}

void lu(double x[][N+1], double l[][N], double u[][N], double y[N]) {
    for (int i=0;i<N;i++) {
        pivot(x, i);
        l[i][i] = x[i][i];
        for (int j=i;j<N;j++) {
            u[i][j] = x[i][j];
        }
        for (int a=i+1;a<N;a++) {
            l[a][i] = x[a][i];
            double det = x[a][i] / x[i][i];
            for (int b=i;b<N+1;b++) {
                x[a][b] -= det * x[i][b];
            }
        }
    }
    for (int i=0;i<N;i++) {
        y[i] = x[i][N];
    }
}

void sol(double x[N], double u[][N], double y[N]) {
    for (int i=N-1;i>=0;i--) {
        x[i] = y[i];
        for (int j=i+1;j<N;j++) {
            x[i] -= u[i][j] * x[j];
        }
        x[i] /= u[i][i];
    }
    printf("x:\n");
    for (int i=0;i<N;i++) {
        printf("x%d = %f\n", i+1, x[i]);
    }
    printf("\n");
}

void con(double a[][N], double x[N]){
double result[N] = {0};
    for (int i = 0; i < N; i++) {
        for (int j = 0; j < N; j++) {
            result[i] += a[i][j] * x[j];
        }
    }
    for (int i = 0; i < N; i++) { 
    	printf("%f\n", result[i]); 
    } 
	printf("\n");
}

int main(void) {
    double x[][N+1] = {
        {3, 1.5, -6, 4.8, 1.2},
        {1, 1.5, -2, -2.4, 0.6},
        {0, -1.5, -2, -1, -2.4},
        {2, 4, -1.8, -0.6, 0}
    };

    double l[N][N] = {0};
    double u[N][N] = {0};
    double y[N] = {0};
    double s[N] = {0};
    
    double a[4][4];
	for(int i=0;i<4;i++){
		for(int j=0;j<4;j++){
			a[i][j]=x[i][j];
		}
	}

    lu(x, l, u, y);
    printf("U行列:\n");
    for (int i=0;i<N;i++) {
        for (int j=0;j<N;j++) {
            printf("%f ", u[i][j]);
        }
        printf("\n");
    }
    printf("\n");
    
    printf("L行列:\n");
    for (int i=0;i<N;i++) {
        for (int j=0;j<N;j++) {
            printf("%f ", l[i][j]);
        }
        printf("\n");
    }
    printf("\n");
    
    printf("y:\n");
    for (int i=0;i<N;i++){
        printf("y%d = %f\n", i+1, y[i]);
    }
    printf("\n");
    
    sol(s, u, y);
    con(a, s);
    return 0;
}