Actual source code: fnutil.c

  1: /*
  2:    - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
  3:    SLEPc - Scalable Library for Eigenvalue Problem Computations
  4:    Copyright (c) 2002-, Universitat Politecnica de Valencia, Spain

  6:    This file is part of SLEPc.
  7:    SLEPc is distributed under a 2-clause BSD license (see LICENSE).
  8:    - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
  9: */
 10: /*
 11:    Utility subroutines common to several impls
 12: */

 14: #include <slepc/private/fnimpl.h>
 15: #include <slepcblaslapack.h>

 17: /*
 18:    Compute the square root of an upper quasi-triangular matrix T,
 19:    using Higham's algorithm (LAA 88, 1987). T is overwritten with sqrtm(T).
 20:  */
 21: static PetscErrorCode SlepcMatDenseSqrt(PetscBLASInt n,PetscScalar *T,PetscBLASInt ld)
 22: {
 23:   PetscScalar  one=1.0,mone=-1.0;
 24:   PetscReal    scal;
 25:   PetscBLASInt i,j,si,sj,r,ione=1;
 26: #if !PetscDefined(USE_COMPLEX)
 27:   PetscReal    alpha,theta,mu,mu2;
 28: #endif

 30:   PetscFunctionBegin;
 31:   for (j=0;j<n;j++) {
 32: #if PetscDefined(USE_COMPLEX)
 33:     sj = 1;
 34:     T[j+j*ld] = PetscSqrtScalar(T[j+j*ld]);
 35: #else
 36:     sj = (j==n-1 || T[j+1+j*ld] == 0.0)? 1: 2;
 37:     if (sj==1) {
 38:       PetscCheck(T[j+j*ld]>=0.0,PETSC_COMM_SELF,PETSC_ERR_USER_INPUT,"Matrix has a real negative eigenvalue, no real primary square root exists");
 39:       T[j+j*ld] = PetscSqrtReal(T[j+j*ld]);
 40:     } else {
 41:       /* square root of 2x2 block */
 42:       theta = (T[j+j*ld]+T[j+1+(j+1)*ld])/2.0;
 43:       mu    = (T[j+j*ld]-T[j+1+(j+1)*ld])/2.0;
 44:       mu2   = -mu*mu-T[j+1+j*ld]*T[j+(j+1)*ld];
 45:       mu    = PetscSqrtReal(mu2);
 46:       if (theta>0.0) alpha = PetscSqrtReal((theta+PetscSqrtReal(theta*theta+mu2))/2.0);
 47:       else alpha = mu/PetscSqrtReal(2.0*(-theta+PetscSqrtReal(theta*theta+mu2)));
 48:       T[j+j*ld]       /= 2.0*alpha;
 49:       T[j+1+(j+1)*ld] /= 2.0*alpha;
 50:       T[j+(j+1)*ld]   /= 2.0*alpha;
 51:       T[j+1+j*ld]     /= 2.0*alpha;
 52:       T[j+j*ld]       += alpha-theta/(2.0*alpha);
 53:       T[j+1+(j+1)*ld] += alpha-theta/(2.0*alpha);
 54:     }
 55: #endif
 56:     for (i=j-1;i>=0;i--) {
 57: #if PetscDefined(USE_COMPLEX)
 58:       si = 1;
 59: #else
 60:       si = (i==0 || T[i+(i-1)*ld] == 0.0)? 1: 2;
 61:       if (si==2) i--;
 62: #endif
 63:       /* solve Sylvester equation of order si x sj */
 64:       r = j-i-si;
 65:       if (r) PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&si,&sj,&r,&mone,T+i+(i+si)*ld,&ld,T+i+si+j*ld,&ld,&one,T+i+j*ld,&ld));
 66:       PetscCallLAPACKInfo("LAPACKtrsyl",LAPACKtrsyl_("N","N",&ione,&si,&sj,T+i+i*ld,&ld,T+j+j*ld,&ld,T+i+j*ld,&ld,&scal,&info));
 67:       PetscCheck(scal==1.0,PETSC_COMM_SELF,PETSC_ERR_SUP,"Current implementation cannot handle scale factor %g",(double)scal);
 68:     }
 69:     if (sj==2) j++;
 70:   }
 71:   PetscFunctionReturn(PETSC_SUCCESS);
 72: }

 74: #define BLOCKSIZE 64

 76: /*
 77:    Schur method for the square root of an upper quasi-triangular matrix T.
 78:    T is overwritten with sqrtm(T).
 79:    If firstonly then only the first column of T will contain relevant values.
 80:  */
 81: PetscErrorCode FNSqrtmSchur(FN fn,PetscBLASInt n,PetscScalar *T,PetscBLASInt ld,PetscBool firstonly)
 82: {
 83:   PetscBLASInt   i,j,k,r,ione=1,sdim,lwork,*s,*p,bs=BLOCKSIZE;
 84:   PetscScalar    *wr,*W,*Q,*work,one=1.0,zero=0.0,mone=-1.0;
 85:   PetscInt       m,nblk;
 86:   PetscReal      scal;
 87: #if PetscDefined(USE_COMPLEX)
 88:   PetscReal      *rwork;
 89: #else
 90:   PetscReal      *wi;
 91: #endif

 93:   PetscFunctionBegin;
 94:   m     = n;
 95:   nblk  = (m+bs-1)/bs;
 96:   lwork = 5*n;
 97:   k     = firstonly? 1: n;

 99:   /* compute Schur decomposition A*Q = Q*T */
100: #if !PetscDefined(USE_COMPLEX)
101:   PetscCall(PetscMalloc7(m,&wr,m,&wi,m*k,&W,m*m,&Q,lwork,&work,nblk,&s,nblk,&p));
102:   PetscCallLAPACKInfo("LAPACKgees",LAPACKgees_("V","N",NULL,&n,T,&ld,&sdim,wr,wi,Q,&ld,work,&lwork,NULL,&info));
103: #else
104:   PetscCall(PetscMalloc7(m,&wr,m,&rwork,m*k,&W,m*m,&Q,lwork,&work,nblk,&s,nblk,&p));
105:   PetscCallLAPACKInfo("LAPACKgees",LAPACKgees_("V","N",NULL,&n,T,&ld,&sdim,wr,Q,&ld,work,&lwork,rwork,NULL,&info));
106: #endif

108:   /* determine block sizes and positions, to avoid cutting 2x2 blocks */
109:   j = 0;
110:   p[j] = 0;
111:   do {
112:     s[j] = PetscMin(bs,n-p[j]);
113: #if !PetscDefined(USE_COMPLEX)
114:     if (p[j]+s[j]!=n && T[p[j]+s[j]+(p[j]+s[j]-1)*ld]!=0.0) s[j]++;
115: #endif
116:     if (p[j]+s[j]==n) break;
117:     j++;
118:     p[j] = p[j-1]+s[j-1];
119:   } while (1);
120:   nblk = j+1;

122:   for (j=0;j<nblk;j++) {
123:     /* evaluate f(T_jj) */
124:     PetscCall(SlepcMatDenseSqrt(s[j],T+p[j]+p[j]*ld,ld));
125:     for (i=j-1;i>=0;i--) {
126:       /* solve Sylvester equation for block (i,j) */
127:       r = p[j]-p[i]-s[i];
128:       if (r) PetscCallBLAS("BLASgemm",BLASgemm_("N","N",s+i,s+j,&r,&mone,T+p[i]+(p[i]+s[i])*ld,&ld,T+p[i]+s[i]+p[j]*ld,&ld,&one,T+p[i]+p[j]*ld,&ld));
129:       PetscCallLAPACKInfo("LAPACKtrsyl",LAPACKtrsyl_("N","N",&ione,s+i,s+j,T+p[i]+p[i]*ld,&ld,T+p[j]+p[j]*ld,&ld,T+p[i]+p[j]*ld,&ld,&scal,&info));
130:       PetscCheck(scal==1.0,PETSC_COMM_SELF,PETSC_ERR_SUP,"Current implementation cannot handle scale factor %g",(double)scal);
131:     }
132:   }

134:   /* backtransform B = Q*T*Q' */
135:   PetscCallBLAS("BLASgemm",BLASgemm_("N","C",&n,&k,&n,&one,T,&ld,Q,&ld,&zero,W,&ld));
136:   PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&k,&n,&one,Q,&ld,W,&ld,&zero,T,&ld));

138:   /* flop count: Schur decomposition, triangular square root, and backtransform */
139:   PetscCall(PetscLogFlops(25.0*n*n*n+n*n*n/3.0+4.0*n*n*k));

141: #if !PetscDefined(USE_COMPLEX)
142:   PetscCall(PetscFree7(wr,wi,W,Q,work,s,p));
143: #else
144:   PetscCall(PetscFree7(wr,rwork,W,Q,work,s,p));
145: #endif
146:   PetscFunctionReturn(PETSC_SUCCESS);
147: }

149: #define DBMAXIT 25

151: /*
152:    Computes the principal square root of the matrix T using the product form
153:    of the Denman-Beavers iteration.
154:    T is overwritten with sqrtm(T) or inv(sqrtm(T)) depending on flag inv.
155:  */
156: PetscErrorCode FNSqrtmDenmanBeavers(FN fn,PetscBLASInt n,PetscScalar *T,PetscBLASInt ld,PetscBool inv)
157: {
158:   PetscScalar        *Told,*M=NULL,*invM,*work,work1,prod,alpha;
159:   PetscScalar        szero=0.0,sone=1.0,smone=-1.0,spfive=0.5,sp25=0.25;
160:   PetscReal          tol,Mres=0.0,detM,g,reldiff,fnormdiff,fnormT,rwork[1];
161:   PetscBLASInt       N,i,it,*piv=NULL,query=-1,lwork;
162:   const PetscBLASInt one=1;
163:   PetscBool          converged=PETSC_FALSE,scale;
164:   unsigned int       ftz;

166:   PetscFunctionBegin;
167:   N = n*n;
168:   tol = PetscSqrtReal((PetscReal)n)*PETSC_MACHINE_EPSILON/2;
169:   scale = PetscDefined(USE_REAL_SINGLE)? PETSC_FALSE: PETSC_TRUE;
170:   PetscCall(SlepcSetFlushToZero(&ftz));

172:   /* query work size */
173:   PetscCallLAPACKInfo("LAPACKgetri",LAPACKgetri_(&n,M,&ld,piv,&work1,&query,&info));
174:   PetscCall(PetscBLASIntCast((PetscInt)PetscRealPart(work1),&lwork));
175:   PetscCall(PetscMalloc5(lwork,&work,n,&piv,n*n,&Told,n*n,&M,n*n,&invM));
176:   PetscCall(PetscArraycpy(M,T,n*n));

178:   if (inv) {  /* start recurrence with I instead of A */
179:     PetscCall(PetscArrayzero(T,n*n));
180:     for (i=0;i<n;i++) T[i+i*ld] += 1.0;
181:   }

183:   for (it=0;it<DBMAXIT && !converged;it++) {

185:     if (scale) {  /* g = (abs(det(M)))^(-1/(2*n)) */
186:       PetscCall(PetscArraycpy(invM,M,n*n));
187:       PetscCallLAPACKInfo("LAPACKgetrf",LAPACKgetrf_(&n,&n,invM,&ld,piv,&info));
188:       prod = invM[0];
189:       for (i=1;i<n;i++) prod *= invM[i+i*ld];
190:       detM = PetscAbsScalar(prod);
191:       g = (detM>PETSC_MAX_REAL)? 0.5: PetscPowReal(detM,-1.0/(2.0*n));
192:       alpha = g;
193:       PetscCallBLAS("BLASscal",BLASscal_(&N,&alpha,T,&one));
194:       alpha = g*g;
195:       PetscCallBLAS("BLASscal",BLASscal_(&N,&alpha,M,&one));
196:       PetscCall(PetscLogFlops(2.0*n*n*n/3.0+2.0*n*n));
197:     }

199:     PetscCall(PetscArraycpy(Told,T,n*n));
200:     PetscCall(PetscArraycpy(invM,M,n*n));

202:     PetscCallLAPACKInfo("LAPACKgetrf",LAPACKgetrf_(&n,&n,invM,&ld,piv,&info));
203:     PetscCallLAPACKInfo("LAPACKgetri",LAPACKgetri_(&n,invM,&ld,piv,work,&lwork,&info));
204:     PetscCall(PetscLogFlops(2.0*n*n*n/3.0+4.0*n*n*n/3.0));

206:     for (i=0;i<n;i++) invM[i+i*ld] += 1.0;
207:     PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&n,&n,&spfive,Told,&ld,invM,&ld,&szero,T,&ld));
208:     for (i=0;i<n;i++) invM[i+i*ld] -= 1.0;

210:     PetscCallBLAS("BLASaxpy",BLASaxpy_(&N,&sone,invM,&one,M,&one));
211:     PetscCallBLAS("BLASscal",BLASscal_(&N,&sp25,M,&one));
212:     for (i=0;i<n;i++) M[i+i*ld] -= 0.5;
213:     PetscCall(PetscLogFlops(2.0*n*n*n+2.0*n*n));

215:     Mres = LAPACKlange_("F",&n,&n,M,&n,rwork);
216:     for (i=0;i<n;i++) M[i+i*ld] += 1.0;

218:     if (scale) {
219:       /* reldiff = norm(T - Told,'fro')/norm(T,'fro') */
220:       PetscCallBLAS("BLASaxpy",BLASaxpy_(&N,&smone,T,&one,Told,&one));
221:       fnormdiff = LAPACKlange_("F",&n,&n,Told,&n,rwork);
222:       fnormT = LAPACKlange_("F",&n,&n,T,&n,rwork);
223:       PetscCall(PetscLogFlops(7.0*n*n));
224:       reldiff = fnormdiff/fnormT;
225:       PetscCall(PetscInfo(fn,"it: %" PetscBLASInt_FMT " reldiff: %g scale: %g tol*scale: %g\n",it,(double)reldiff,(double)g,(double)(tol*g)));
226:       if (reldiff<1e-2) scale = PETSC_FALSE;  /* Switch off scaling */
227:     }

229:     if (Mres<=tol) converged = PETSC_TRUE;
230:   }

232:   PetscCheck(Mres<=tol,PETSC_COMM_SELF,PETSC_ERR_LIB,"SQRTM not converged after %d iterations",DBMAXIT);
233:   PetscCall(PetscFree5(work,piv,Told,M,invM));
234:   PetscCall(SlepcResetFlushToZero(&ftz));
235:   PetscFunctionReturn(PETSC_SUCCESS);
236: }

238: #define NSMAXIT 50

240: /*
241:    Computes the principal square root of the matrix A using the Newton-Schulz iteration.
242:    T is overwritten with sqrtm(T) or inv(sqrtm(T)) depending on flag inv.
243:  */
244: PetscErrorCode FNSqrtmNewtonSchulz(FN fn,PetscBLASInt n,PetscScalar *A,PetscBLASInt ld,PetscBool inv)
245: {
246:   PetscScalar    *Y=A,*Yold,*Z,*Zold,*M;
247:   PetscScalar    szero=0.0,sone=1.0,smone=-1.0,spfive=0.5,sthree=3.0;
248:   PetscReal      sqrtnrm,tol,Yres=0.0,nrm,rwork[1],done=1.0;
249:   PetscBLASInt   i,it,N,one=1,zero=0;
250:   PetscBool      converged=PETSC_FALSE;
251:   unsigned int   ftz;

253:   PetscFunctionBegin;
254:   N = n*n;
255:   tol = PetscSqrtReal((PetscReal)n)*PETSC_MACHINE_EPSILON/2;
256:   PetscCall(SlepcSetFlushToZero(&ftz));

258:   PetscCall(PetscMalloc4(N,&Yold,N,&Z,N,&Zold,N,&M));

260:   /* scale */
261:   PetscCall(PetscArraycpy(Z,A,N));
262:   for (i=0;i<n;i++) Z[i+i*ld] -= 1.0;
263:   nrm = LAPACKlange_("fro",&n,&n,Z,&n,rwork);
264:   sqrtnrm = PetscSqrtReal(nrm);
265:   PetscCallLAPACKInfo("LAPACKlascl",LAPACKlascl_("G",&zero,&zero,&nrm,&done,&N,&one,A,&N,&info));
266:   tol *= nrm;
267:   PetscCall(PetscInfo(fn,"||I-A||_F = %g, new tol: %g\n",(double)nrm,(double)tol));
268:   PetscCall(PetscLogFlops(2.0*n*n));

270:   /* Z = I */
271:   PetscCall(PetscArrayzero(Z,N));
272:   for (i=0;i<n;i++) Z[i+i*ld] = 1.0;

274:   for (it=0;it<NSMAXIT && !converged;it++) {
275:     /* Yold = Y, Zold = Z */
276:     PetscCall(PetscArraycpy(Yold,Y,N));
277:     PetscCall(PetscArraycpy(Zold,Z,N));

279:     /* M = (3*I-Zold*Yold) */
280:     PetscCall(PetscArrayzero(M,N));
281:     for (i=0;i<n;i++) M[i+i*ld] = sthree;
282:     PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&n,&n,&smone,Zold,&ld,Yold,&ld,&sone,M,&ld));

284:     /* Y = (1/2)*Yold*M, Z = (1/2)*M*Zold */
285:     PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&n,&n,&spfive,Yold,&ld,M,&ld,&szero,Y,&ld));
286:     PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&n,&n,&spfive,M,&ld,Zold,&ld,&szero,Z,&ld));

288:     /* reldiff = norm(Y-Yold,'fro')/norm(Y,'fro') */
289:     PetscCallBLAS("BLASaxpy",BLASaxpy_(&N,&smone,Y,&one,Yold,&one));
290:     Yres = LAPACKlange_("fro",&n,&n,Yold,&n,rwork);
291:     PetscCheck(!PetscIsNanReal(Yres),PETSC_COMM_SELF,PETSC_ERR_FP,"The computed norm is not-a-number");
292:     if (Yres<=tol) converged = PETSC_TRUE;
293:     PetscCall(PetscInfo(fn,"it: %" PetscBLASInt_FMT " res: %g\n",it,(double)Yres));

295:     PetscCall(PetscLogFlops(6.0*n*n*n+2.0*n*n));
296:   }

298:   PetscCheck(Yres<=tol,PETSC_COMM_SELF,PETSC_ERR_LIB,"SQRTM not converged after %d iterations",NSMAXIT);

300:   /* undo scaling */
301:   if (inv) {
302:     PetscCall(PetscArraycpy(A,Z,N));
303:     PetscCallLAPACKInfo("LAPACKlascl",LAPACKlascl_("G",&zero,&zero,&sqrtnrm,&done,&N,&one,A,&N,&info));
304:   } else PetscCallLAPACKInfo("LAPACKlascl",LAPACKlascl_("G",&zero,&zero,&done,&sqrtnrm,&N,&one,A,&N,&info));

306:   PetscCall(PetscFree4(Yold,Z,Zold,M));
307:   PetscCall(SlepcResetFlushToZero(&ftz));
308:   PetscFunctionReturn(PETSC_SUCCESS);
309: }

311: #if PetscDefined(HAVE_CUDA)
312: #include "../src/sys/classes/fn/impls/cuda/fnutilcuda.h"
313: #include <slepccupmblas.h>

315: /*
316:  * Matrix square root by Newton-Schulz iteration. CUDA version.
317:  * Computes the principal square root of the matrix A using the
318:  * Newton-Schulz iteration. A is overwritten with sqrtm(A).
319:  */
320: PetscErrorCode FNSqrtmNewtonSchulz_CUDA(FN fn,PetscBLASInt n,PetscScalar *d_A,PetscBLASInt ld,PetscBool inv)
321: {
322:   PetscScalar        *d_Yold,*d_Z,*d_Zold,*d_M,alpha;
323:   PetscReal          nrm,sqrtnrm,tol,Yres=0.0;
324:   const PetscScalar  szero=0.0,sone=1.0,smone=-1.0,spfive=0.5,sthree=3.0;
325:   PetscInt           it;
326:   PetscBLASInt       N;
327:   const PetscBLASInt one=1;
328:   PetscBool          converged=PETSC_FALSE;
329:   cublasHandle_t     cublasv2handle;

331:   PetscFunctionBegin;
332:   PetscCall(PetscDeviceInitialize(PETSC_DEVICE_CUDA)); /* For CUDA event timers */
333:   PetscCall(PetscCUBLASGetHandle(&cublasv2handle));
334:   N = n*n;
335:   tol = PetscSqrtReal((PetscReal)n)*PETSC_MACHINE_EPSILON/2;

337:   PetscCallCUDA(cudaMalloc((void **)&d_Yold,sizeof(PetscScalar)*N));
338:   PetscCallCUDA(cudaMalloc((void **)&d_Z,sizeof(PetscScalar)*N));
339:   PetscCallCUDA(cudaMalloc((void **)&d_Zold,sizeof(PetscScalar)*N));
340:   PetscCallCUDA(cudaMalloc((void **)&d_M,sizeof(PetscScalar)*N));

342:   PetscCall(PetscLogGpuTimeBegin());

344:   /* Z = I; */
345:   PetscCallCUDA(cudaMemset(d_Z,0,sizeof(PetscScalar)*N));
346:   PetscCall(set_diagonal(n,d_Z,ld,sone));

348:   /* scale */
349:   PetscCallCUBLAS(cublasXaxpy(cublasv2handle,N,&smone,d_A,one,d_Z,one));
350:   PetscCallCUBLAS(cublasXnrm2(cublasv2handle,N,d_Z,one,&nrm));
351:   sqrtnrm = PetscSqrtReal(nrm);
352:   alpha = 1.0/nrm;
353:   PetscCallCUBLAS(cublasXscal(cublasv2handle,N,&alpha,d_A,one));
354:   tol *= nrm;
355:   PetscCall(PetscInfo(fn,"||I-A||_F = %g, new tol: %g\n",(double)nrm,(double)tol));
356:   PetscCall(PetscLogGpuFlops(2.0*n*n));

358:   /* Z = I; */
359:   PetscCallCUDA(cudaMemset(d_Z,0,sizeof(PetscScalar)*N));
360:   PetscCall(set_diagonal(n,d_Z,ld,sone));

362:   for (it=0;it<NSMAXIT && !converged;it++) {
363:     /* Yold = Y, Zold = Z */
364:     PetscCallCUDA(cudaMemcpy(d_Yold,d_A,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));
365:     PetscCallCUDA(cudaMemcpy(d_Zold,d_Z,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));

367:     /* M = (3*I - Zold*Yold) */
368:     PetscCallCUDA(cudaMemset(d_M,0,sizeof(PetscScalar)*N));
369:     PetscCall(set_diagonal(n,d_M,ld,sthree));
370:     PetscCallCUBLAS(cublasXgemm(cublasv2handle,CUBLAS_OP_N,CUBLAS_OP_N,n,n,n,&smone,d_Zold,ld,d_Yold,ld,&sone,d_M,ld));

372:     /* Y = (1/2) * Yold * M, Z = (1/2) * M * Zold */
373:     PetscCallCUBLAS(cublasXgemm(cublasv2handle,CUBLAS_OP_N,CUBLAS_OP_N,n,n,n,&spfive,d_Yold,ld,d_M,ld,&szero,d_A,ld));
374:     PetscCallCUBLAS(cublasXgemm(cublasv2handle,CUBLAS_OP_N,CUBLAS_OP_N,n,n,n,&spfive,d_M,ld,d_Zold,ld,&szero,d_Z,ld));

376:     /* reldiff = norm(Y-Yold,'fro')/norm(Y,'fro') */
377:     PetscCallCUBLAS(cublasXaxpy(cublasv2handle,N,&smone,d_A,one,d_Yold,one));
378:     PetscCallCUBLAS(cublasXnrm2(cublasv2handle,N,d_Yold,one,&Yres));
379:     PetscCheck(!PetscIsNanReal(Yres),PETSC_COMM_SELF,PETSC_ERR_FP,"The computed norm is not-a-number");
380:     if (Yres<=tol) converged = PETSC_TRUE;
381:     PetscCall(PetscInfo(fn,"it: %" PetscInt_FMT " res: %g\n",it,(double)Yres));

383:     PetscCall(PetscLogGpuFlops(6.0*n*n*n+2.0*n*n));
384:   }

386:   PetscCheck(Yres<=tol,PETSC_COMM_SELF,PETSC_ERR_LIB,"SQRTM not converged after %d iterations", NSMAXIT);

388:   /* undo scaling */
389:   if (inv) {
390:     alpha = 1.0/sqrtnrm;
391:     PetscCallCUBLAS(cublasXscal(cublasv2handle,N,&alpha,d_Z,one));
392:     PetscCallCUDA(cudaMemcpy(d_A,d_Z,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));
393:   } else {
394:     alpha = sqrtnrm;
395:     PetscCallCUBLAS(cublasXscal(cublasv2handle,N,&alpha,d_A,one));
396:   }

398:   PetscCall(PetscLogGpuTimeEnd());
399:   PetscCallCUDA(cudaFree(d_Yold));
400:   PetscCallCUDA(cudaFree(d_Z));
401:   PetscCallCUDA(cudaFree(d_Zold));
402:   PetscCallCUDA(cudaFree(d_M));
403:   PetscFunctionReturn(PETSC_SUCCESS);
404: }

406: #if PetscDefined(HAVE_MAGMA)
407: #include <slepcmagma.h>

409: /*
410:  * Matrix square root by product form of Denman-Beavers iteration. CUDA version.
411:  * Computes the principal square root of the matrix T using the product form
412:  * of the Denman-Beavers iteration. T is overwritten with sqrtm(T).
413:  */
414: PetscErrorCode FNSqrtmDenmanBeavers_CUDAm(FN fn,PetscBLASInt n,PetscScalar *d_T,PetscBLASInt ld,PetscBool inv)
415: {
416:   PetscScalar    *d_Told,*d_M,*d_invM,*d_work,prod,szero=0.0,sone=1.0,smone=-1.0,spfive=0.5,sneg_pfive=-0.5,sp25=0.25,alpha;
417:   PetscReal      tol,Mres=0.0,detM,g,reldiff,fnormdiff,fnormT;
418:   PetscInt       it,lwork,nb;
419:   PetscBLASInt   N,one=1,*piv=NULL;
420:   PetscBool      converged=PETSC_FALSE,scale;
421:   cublasHandle_t cublasv2handle;

423:   PetscFunctionBegin;
424:   PetscCall(PetscDeviceInitialize(PETSC_DEVICE_CUDA)); /* For CUDA event timers */
425:   PetscCall(PetscCUBLASGetHandle(&cublasv2handle));
426:   PetscCall(SlepcMagmaInit());
427:   N = n*n;
428:   scale = PetscDefined(USE_REAL_SINGLE)? PETSC_FALSE: PETSC_TRUE;
429:   tol = PetscSqrtReal((PetscReal)n)*PETSC_MACHINE_EPSILON/2;

431:   /* query work size */
432:   nb = magma_get_xgetri_nb(n);
433:   lwork = nb*n;
434:   PetscCall(PetscMalloc1(n,&piv));
435:   PetscCallCUDA(cudaMalloc((void **)&d_work,sizeof(PetscScalar)*lwork));
436:   PetscCallCUDA(cudaMalloc((void **)&d_Told,sizeof(PetscScalar)*N));
437:   PetscCallCUDA(cudaMalloc((void **)&d_M,sizeof(PetscScalar)*N));
438:   PetscCallCUDA(cudaMalloc((void **)&d_invM,sizeof(PetscScalar)*N));

440:   PetscCall(PetscLogGpuTimeBegin());
441:   PetscCallCUDA(cudaMemcpy(d_M,d_T,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));
442:   if (inv) {  /* start recurrence with I instead of A */
443:     PetscCallCUDA(cudaMemset(d_T,0,sizeof(PetscScalar)*N));
444:     PetscCall(set_diagonal(n,d_T,ld,1.0));
445:   }

447:   for (it=0;it<DBMAXIT && !converged;it++) {

449:     if (scale) { /* g = (abs(det(M)))^(-1/(2*n)); */
450:       PetscCallCUDA(cudaMemcpy(d_invM,d_M,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));
451:       PetscCallMAGMA(magma_xgetrf_gpu,n,n,d_invM,ld,piv);
452:       PetscCall(mult_diagonal(n,d_invM,ld,&prod));
453:       detM = PetscAbsScalar(prod);
454:       g = (detM>PETSC_MAX_REAL)? 0.5: PetscPowReal(detM,-1.0/(2.0*n));
455:       alpha = g;
456:       PetscCallCUBLAS(cublasXscal(cublasv2handle,N,&alpha,d_T,one));
457:       alpha = g*g;
458:       PetscCallCUBLAS(cublasXscal(cublasv2handle,N,&alpha,d_M,one));
459:       PetscCall(PetscLogGpuFlops(2.0*n*n*n/3.0+2.0*n*n));
460:     }

462:     PetscCallCUDA(cudaMemcpy(d_Told,d_T,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));
463:     PetscCallCUDA(cudaMemcpy(d_invM,d_M,sizeof(PetscScalar)*N,cudaMemcpyDeviceToDevice));

465:     PetscCallMAGMA(magma_xgetrf_gpu,n,n,d_invM,ld,piv);
466:     PetscCallMAGMA(magma_xgetri_gpu,n,d_invM,ld,piv,d_work,lwork);
467:     PetscCall(PetscLogGpuFlops(2.0*n*n*n/3.0+4.0*n*n*n/3.0));

469:     PetscCall(shift_diagonal(n,d_invM,ld,sone));
470:     PetscCallCUBLAS(cublasXgemm(cublasv2handle,CUBLAS_OP_N,CUBLAS_OP_N,n,n,n,&spfive,d_Told,ld,d_invM,ld,&szero,d_T,ld));
471:     PetscCall(shift_diagonal(n,d_invM,ld,smone));

473:     PetscCallCUBLAS(cublasXaxpy(cublasv2handle,N,&sone,d_invM,one,d_M,one));
474:     PetscCallCUBLAS(cublasXscal(cublasv2handle,N,&sp25,d_M,one));
475:     PetscCall(shift_diagonal(n,d_M,ld,sneg_pfive));
476:     PetscCall(PetscLogGpuFlops(2.0*n*n*n+2.0*n*n));

478:     PetscCallCUBLAS(cublasXnrm2(cublasv2handle,N,d_M,one,&Mres));
479:     PetscCall(shift_diagonal(n,d_M,ld,sone));

481:     if (scale) {
482:       /* reldiff = norm(T - Told,'fro')/norm(T,'fro'); */
483:       PetscCallCUBLAS(cublasXaxpy(cublasv2handle,N,&smone,d_T,one,d_Told,one));
484:       PetscCallCUBLAS(cublasXnrm2(cublasv2handle,N,d_Told,one,&fnormdiff));
485:       PetscCallCUBLAS(cublasXnrm2(cublasv2handle,N,d_T,one,&fnormT));
486:       PetscCall(PetscLogGpuFlops(7.0*n*n));
487:       reldiff = fnormdiff/fnormT;
488:       PetscCall(PetscInfo(fn,"it: %" PetscInt_FMT " reldiff: %g scale: %g tol*scale: %g\n",it,(double)reldiff,(double)g,(double)tol*g));
489:       if (reldiff<1e-2) scale = PETSC_FALSE; /* Switch to no scaling. */
490:     }

492:     PetscCall(PetscInfo(fn,"it: %" PetscInt_FMT " Mres: %g\n",it,(double)Mres));
493:     if (Mres<=tol) converged = PETSC_TRUE;
494:   }

496:   PetscCheck(Mres<=tol,PETSC_COMM_SELF,PETSC_ERR_LIB,"SQRTM not converged after %d iterations", DBMAXIT);
497:   PetscCall(PetscLogGpuTimeEnd());
498:   PetscCall(PetscFree(piv));
499:   PetscCallCUDA(cudaFree(d_work));
500:   PetscCallCUDA(cudaFree(d_Told));
501:   PetscCallCUDA(cudaFree(d_M));
502:   PetscCallCUDA(cudaFree(d_invM));
503:   PetscFunctionReturn(PETSC_SUCCESS);
504: }
505: #endif /* PETSC_HAVE_MAGMA */

507: #endif /* PETSC_HAVE_CUDA */

509: #define ITMAX 5

511: /*
512:    Estimate norm(A^m,1) by block 1-norm power method (required workspace is 11*n)
513: */
514: static PetscErrorCode SlepcNormEst1(PetscBLASInt n,PetscScalar *A,PetscInt m,PetscScalar *work,PetscRandom rand,PetscReal *nrm)
515: {
516:   PetscScalar    *X,*Y,*Z,*S,*S_old,*aux,val,sone=1.0,szero=0.0;
517:   PetscReal      est=0.0,est_old,vals[2]={0.0,0.0},*zvals,maxzval[2],raux;
518:   PetscBLASInt   i,j,t=2,it=0,ind[2],est_j=0,m1;

520:   PetscFunctionBegin;
521:   X = work;
522:   Y = work + 2*n;
523:   Z = work + 4*n;
524:   S = work + 6*n;
525:   S_old = work + 8*n;
526:   zvals = (PetscReal*)(work + 10*n);

528:   for (i=0;i<n;i++) {  /* X has columns of unit 1-norm */
529:     X[i] = 1.0/n;
530:     PetscCall(PetscRandomGetValue(rand,&val));
531:     if (PetscRealPart(val) < 0.5) X[i+n] = -1.0/n;
532:     else X[i+n] = 1.0/n;
533:   }
534:   for (i=0;i<t*n;i++) S[i] = 0.0;
535:   ind[0] = 0; ind[1] = 0;
536:   est_old = 0;
537:   while (1) {
538:     it++;
539:     for (j=0;j<m;j++) {  /* Y = A^m*X */
540:       PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&t,&n,&sone,A,&n,X,&n,&szero,Y,&n));
541:       if (j<m-1) SlepcSwap(X,Y,aux);
542:     }
543:     for (j=0;j<t;j++) {  /* vals[j] = norm(Y(:,j),1) */
544:       vals[j] = 0.0;
545:       for (i=0;i<n;i++) vals[j] += PetscAbsScalar(Y[i+j*n]);
546:     }
547:     if (vals[0]<vals[1]) {
548:       SlepcSwap(vals[0],vals[1],raux);
549:       m1 = 1;
550:     } else m1 = 0;
551:     est = vals[0];
552:     if (est>est_old || it==2) est_j = ind[m1];
553:     if (it>=2 && est<=est_old) {
554:       est = est_old;
555:       break;
556:     }
557:     est_old = est;
558:     if (it>ITMAX) break;
559:     SlepcSwap(S,S_old,aux);
560:     for (i=0;i<t*n;i++) {  /* S = sign(Y) */
561:       S[i] = (PetscRealPart(Y[i]) < 0.0)? -1.0: 1.0;
562:     }
563:     for (j=0;j<m;j++) {  /* Z = (A^T)^m*S */
564:       PetscCallBLAS("BLASgemm",BLASgemm_("C","N",&n,&t,&n,&sone,A,&n,S,&n,&szero,Z,&n));
565:       if (j<m-1) SlepcSwap(S,Z,aux);
566:     }
567:     maxzval[0] = -1; maxzval[1] = -1;
568:     ind[0] = 0; ind[1] = 0;
569:     for (i=0;i<n;i++) {  /* zvals[i] = norm(Z(i,:),inf) */
570:       zvals[i] = PetscMax(PetscAbsScalar(Z[i+0*n]),PetscAbsScalar(Z[i+1*n]));
571:       if (zvals[i]>maxzval[0]) {
572:         maxzval[0] = zvals[i];
573:         ind[0] = i;
574:       } else if (zvals[i]>maxzval[1]) {
575:         maxzval[1] = zvals[i];
576:         ind[1] = i;
577:       }
578:     }
579:     if (it>=2 && maxzval[0]==zvals[est_j]) break;
580:     for (i=0;i<t*n;i++) X[i] = 0.0;
581:     for (j=0;j<t;j++) X[ind[j]+j*n] = 1.0;
582:   }
583:   *nrm = est;
584:   /* Flop count is roughly (it * 2*m * t*gemv) = 4*its*m*t*n*n */
585:   PetscCall(PetscLogFlops(4.0*it*m*t*n*n));
586:   PetscFunctionReturn(PETSC_SUCCESS);
587: }

589: #define SMALLN 100

591: /*
592:    Estimate norm(A^m,1) (required workspace is 2*n*n)
593: */
594: PetscErrorCode SlepcNormAm(PetscBLASInt n,PetscScalar *A,PetscInt m,PetscScalar *work,PetscRandom rand,PetscReal *nrm)
595: {
596:   PetscScalar    *v=work,*w=work+n*n,*aux,sone=1.0,szero=0.0;
597:   PetscReal      rwork[1],tmp;
598:   PetscBLASInt   i,j,one=1;
599:   PetscBool      isrealpos=PETSC_TRUE;

601:   PetscFunctionBegin;
602:   if (n<SMALLN) {   /* compute matrix power explicitly */
603:     if (m==1) {
604:       *nrm = LAPACKlange_("O",&n,&n,A,&n,rwork);
605:       PetscCall(PetscLogFlops(1.0*n*n));
606:     } else {  /* m>=2 */
607:       PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&n,&n,&sone,A,&n,A,&n,&szero,v,&n));
608:       for (j=0;j<m-2;j++) {
609:         PetscCallBLAS("BLASgemm",BLASgemm_("N","N",&n,&n,&n,&sone,A,&n,v,&n,&szero,w,&n));
610:         SlepcSwap(v,w,aux);
611:       }
612:       *nrm = LAPACKlange_("O",&n,&n,v,&n,rwork);
613:       PetscCall(PetscLogFlops(2.0*n*n*n*(m-1)+1.0*n*n));
614:     }
615:   } else {
616:     for (i=0;i<n;i++)
617:       for (j=0;j<n;j++)
618: #if PetscDefined(USE_COMPLEX)
619:         if (PetscRealPart(A[i+j*n])<0.0 || PetscImaginaryPart(A[i+j*n])!=0.0) { isrealpos = PETSC_FALSE; break; }
620: #else
621:         if (A[i+j*n]<0.0) { isrealpos = PETSC_FALSE; break; }
622: #endif
623:     if (isrealpos) {   /* for positive matrices only */
624:       for (i=0;i<n;i++) v[i] = 1.0;
625:       for (j=0;j<m;j++) {  /* w = A'*v */
626:         PetscCallBLAS("BLASgemv",BLASgemv_("C",&n,&n,&sone,A,&n,v,&one,&szero,w,&one));
627:         SlepcSwap(v,w,aux);
628:       }
629:       PetscCall(PetscLogFlops(2.0*n*n*m));
630:       *nrm = 0.0;
631:       for (i=0;i<n;i++) if ((tmp = PetscAbsScalar(v[i])) > *nrm) *nrm = tmp;   /* norm(v,inf) */
632:     } else PetscCall(SlepcNormEst1(n,A,m,work,rand,nrm));
633:   }
634:   PetscFunctionReturn(PETSC_SUCCESS);
635: }